diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 50fd2e9e5..c4ee02cab 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2016,7 +2016,7 @@ class Scheduler( draft_worker=self.draft_worker, model_worker=self.model_worker, logprob_result_processor=SchedulerLogprobResultProcessor( - server_args=self.server_args, model_config=self.model_config + model_config=self.model_config ), output_streamer=self.output_streamer, abort_request=self.abort_request, diff --git a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py index f1f659565..379501f0e 100644 --- a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py @@ -14,17 +14,13 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.io_struct import build_flat_input_top_logprobs_arrays from sglang.srt.managers.schedule_batch import Req from sglang.srt.runtime_context import get_exec -from sglang.srt.server_args import ( - MIS_DELIMITER_TOKEN_ID, - ServerArgs, -) +from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID logger = logging.getLogger(__name__) @dataclass(kw_only=True, slots=True, frozen=True) class SchedulerLogprobResultProcessor: - server_args: ServerArgs model_config: ModelConfig def _process_input_token_logprobs( diff --git a/test/registered/unit/managers/test_flat_raw_top_logprobs.py b/test/registered/unit/managers/test_flat_raw_top_logprobs.py index d4225dd9a..a32e26d49 100644 --- a/test/registered/unit/managers/test_flat_raw_top_logprobs.py +++ b/test/registered/unit/managers/test_flat_raw_top_logprobs.py @@ -19,6 +19,7 @@ from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel maybe_stub_sgl_kernel() +from sglang.srt import runtime_context as rc from sglang.srt.managers.io_struct import ( BatchTokenIDOutput, GenerateReqInput, @@ -318,9 +319,8 @@ class TestB64MetaInfo(CustomTestCase): def _make_logprob_processor() -> SchedulerLogprobResultProcessor: - # The processor only reads enable_mis and vocab_size from these. + # enable_mis comes from the published exec bag, not from here; see setUp. return SchedulerLogprobResultProcessor( - server_args=SimpleNamespace(enable_mis=False), model_config=SimpleNamespace(vocab_size=1_000_000), ) @@ -341,6 +341,14 @@ _SCHED_IDX_ROWS = [[11, 22], [33, 44], [55, 66], [77, 88], [99, 100]] class TestSchedulerFlatAssembly(CustomTestCase): """Scheduler-side flat assembly in the logprob result processor.""" + def setUp(self): + super().setUp() + self._server_args_override = rc.get_context().override_server_args() + self._server_args_override.install() + + def tearDown(self): + self._server_args_override.restore() + def _make_req(self, flat: bool, num_tokens: int = 5) -> Req: return Req( "r0",