config: retire the hidden global fallbacks and the mamba-extra-buffer instance reads (#34080)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user