config: retire the hidden global fallbacks and the mamba-extra-buffer instance reads (#34080)

This commit is contained in:
Cheng Wan
2026-08-09 14:43:28 -07:00
committed by GitHub
parent 967cac801c
commit 110bf7e6a8
12 changed files with 112 additions and 83 deletions
@@ -1232,7 +1232,6 @@ class TestMlxOverlapScheduler(unittest.TestCase):
logits_output = SimpleNamespace(customized_info=None)
original_release = batch_result_processor_module.release_kv_cache
original_get_indexer = batch_result_processor_module.get_global_indexer_capturer
original_get_server_args = batch_result_processor_module.get_server_args
def fake_release_kv_cache(release_req, tree_cache, is_insert=False):
events.append(("release", release_req.rid))
@@ -1240,21 +1239,26 @@ class TestMlxOverlapScheduler(unittest.TestCase):
batch_result_processor_module.release_kv_cache = fake_release_kv_cache
batch_result_processor_module.get_global_indexer_capturer = lambda: None
batch_result_processor_module.get_server_args = lambda: SimpleNamespace(
enable_mamba_extra_buffer_lazy=lambda: False
# The lazy predicate reads the published bags; publish the non-lazy
# strategy instead of stubbing the accessor.
from sglang.srt.runtime_context import get_context
override = get_context().override_server_args(
mamba_radix_cache_strategy="extra_buffer"
)
override.install()
try:
SchedulerBatchResultProcessor._handle_finish_state_updated_req(
processor, req, batch, result, i, logits_output
)
finally:
override.restore()
for name, original in saved.items():
setattr(SchedulerBatchResultProcessor, name, original)
batch_result_processor_module.release_kv_cache = original_release
batch_result_processor_module.get_global_indexer_capturer = (
original_get_indexer
)
batch_result_processor_module.get_server_args = original_get_server_args
self.assertEqual(
events,
@@ -10,6 +10,7 @@ from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.runtime_context import get_context
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.test.ci.ci_register import register_cpu_ci
@@ -144,42 +145,25 @@ class TestMambaBoundaryMaskReuse(unittest.TestCase):
processor.process_batch_result_decode(result_batch, batch_result)
scheduler.process_batch_result = process_batch_result
server_args = SimpleNamespace(
enable_mamba_extra_buffer=lambda: True,
enable_mamba_extra_buffer_lazy=lambda: False,
)
with (
# The mamba predicates and the track interval read the
# published bags, so publish the configuration under test
# (non-lazy extra buffer, interval 4); observability and
# disagg reads are served by the same publish at their
# defaults.
get_context().override_server_args(
mamba_radix_cache_strategy="extra_buffer",
mamba_track_interval=4,
),
patch(
"sglang.srt.managers.schedule_batch.alloc_for_decode",
return_value=torch.tensor([3], dtype=torch.int64),
),
patch(
"sglang.srt.managers.schedule_batch.get_server_args",
return_value=server_args,
),
patch(
"sglang.srt.managers.schedule_batch.get_exec",
return_value=SimpleNamespace(
mamba=SimpleNamespace(mamba_track_interval=4)
),
),
patch(
"sglang.srt.managers.schedule_batch.set_mamba_track_indices_from_reqs"
),
patch.object(torch.Tensor, "pin_memory", lambda tensor: tensor),
patch(
"sglang.srt.managers.scheduler_components."
"batch_result_processor.get_observability",
return_value=SimpleNamespace(enable_metrics=False),
),
patch(
"sglang.srt.managers.scheduler_components."
"batch_result_processor.get_disagg",
return_value=SimpleNamespace(
disaggregation_decode_enable_offload_kvcache=False
),
),
patch.object(
SchedulerBatchResultProcessor,
"_mamba_prefix_cache_update",
@@ -4,6 +4,7 @@ from unittest.mock import patch
import torch
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import maybe_stub_sgl_kernel
@@ -43,18 +44,16 @@ class TestPrepareForDecodeSeqLensOwnership(unittest.TestCase):
"""Each prepare_for_decode call rebinds seq-lens tensors to new +1 objects without mutating the old ones."""
batch = _make_decode_batch()
server_args = types.SimpleNamespace(
enable_mamba_extra_buffer=lambda: False,
# The mamba-extra-buffer predicate reads the published bags, so the
# fixture publishes a config with the strategy off.
override = get_context().override_server_args(
mamba_radix_cache_strategy="no_buffer"
)
with (
patch(
"sglang.srt.managers.schedule_batch.alloc_for_decode",
return_value=torch.tensor([6, 7], dtype=torch.int64),
),
patch(
"sglang.srt.managers.schedule_batch.get_server_args",
return_value=server_args,
),
override.install()
self.addCleanup(override.restore)
with patch(
"sglang.srt.managers.schedule_batch.alloc_for_decode",
return_value=torch.tensor([6, 7], dtype=torch.int64),
):
for step in range(1, 3):
prev_seq_lens = batch.seq_lens