From 981dfa2b83badaf6a8a6b2c14bbf24d7e267bb23 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Mon, 24 Aug 2026 20:17:50 +0800 Subject: [PATCH] Make the scheduler track the published weight version (#35925) --- python/sglang/srt/entrypoints/http_server.py | 4 +- python/sglang/srt/managers/io_struct.py | 4 + python/sglang/srt/managers/scheduler.py | 20 ++ .../scheduler_components/weight_updater.py | 13 +- .../srt/managers/tokenizer_control_mixin.py | 10 + .../test_scheduler_weight_version_tracking.py | 187 ++++++++++++++++++ 6 files changed, 233 insertions(+), 5 deletions(-) create mode 100644 test/registered/unit/managers/test_scheduler_weight_version_tracking.py diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index e92a8eaa0..cb363cca7 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -1442,9 +1442,7 @@ async def update_weight_version( # Use a simple approach without the complex lock mechanism for now # since weight_version update is a simple operation that doesn't affect model weights try: - _global_state.tokenizer_manager.record_config_updates( - "http.update_weight_version", weight_version=obj.new_version - ) + await _global_state.tokenizer_manager.update_weight_version(obj) return ORJSONResponse( { diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index f9f2f7df2..7e4e6e7f0 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1923,6 +1923,10 @@ class UpdateWeightVersionReqInput(BaseReq, kw_only=True): abort_all_requests: bool = True +class UpdateWeightVersionReqOutput(BaseReq, kw_only=True): + pass + + class GetWeightsByNameReqInput(BaseReq, kw_only=True): name: str truncate_size: int = 100 diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 4f6032fa6..571fde8d0 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -179,6 +179,8 @@ from sglang.srt.managers.io_struct import ( UpdateWeightsFromDistributedReqInput, UpdateWeightsFromIPCReqInput, UpdateWeightsFromTensorReqInput, + UpdateWeightVersionReqInput, + UpdateWeightVersionReqOutput, sock_send, ) from sglang.srt.managers.load_snapshot import create_load_snapshot_writer @@ -1608,6 +1610,10 @@ class Scheduler( UpdateWeightsFromIPCReqInput, self.weight_updater.update_weights_from_ipc, ), + ( + UpdateWeightVersionReqInput, + self.handle_update_weight_version, + ), ( GetWeightsByNameReqInput, self.weight_updater.get_weights_by_name, @@ -4568,6 +4574,20 @@ class Scheduler( barrier(group=self.tp_group.cpu_group) return RpcReqOutput(success=success, message="" if not exec else str(exec)) + def handle_update_weight_version( + self, recv_req: UpdateWeightVersionReqInput + ) -> UpdateWeightVersionReqOutput: + self.record_weight_version_change(new_version=recv_req.new_version) + return UpdateWeightVersionReqOutput() + + def record_weight_version_change(self, new_version: Optional[str]) -> None: + if new_version is None or new_version == get_serving().weight_version: + return + + old_version = get_serving().weight_version + get_context().override("scheduler.weight_version", weight_version=new_version) + logger.info(f"Weight version changed. {old_version=} {new_version=}") + def collect_inflight_reqs(self) -> Set[Req]: if self.ps.pp_size == 1: inflight_batches = [self.running_batch, self.last_batch] diff --git a/python/sglang/srt/managers/scheduler_components/weight_updater.py b/python/sglang/srt/managers/scheduler_components/weight_updater.py index 653b28c41..362de6889 100644 --- a/python/sglang/srt/managers/scheduler_components/weight_updater.py +++ b/python/sglang/srt/managers/scheduler_components/weight_updater.py @@ -107,6 +107,9 @@ class SchedulerWeightUpdaterManager: ) assert flush_cache_success, "Cache flush failed after updating weights" + def record_weight_version_after_update(self, weight_version: Optional[str]) -> None: + self.scheduler.record_weight_version_change(new_version=weight_version) + def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput): """In-place update of the weights from disk.""" with self._observe_weight_load("disk"): @@ -116,7 +119,9 @@ class SchedulerWeightUpdaterManager: success, message = self.draft_worker.update_weights_from_disk(recv_req) if tp_success: self.flush_cache_after_weight_update(recv_req) - if not success: + if success: + self.record_weight_version_after_update(recv_req.weight_version) + else: logger.error(message) return UpdateWeightFromDiskReqOutput( success=success, message=message, num_paused_requests=0 @@ -144,6 +149,7 @@ class SchedulerWeightUpdaterManager: success, message = self.tp_worker.update_weights_from_distributed(recv_req) if success: self.flush_cache_after_weight_update(recv_req) + self.record_weight_version_after_update(recv_req.weight_version) else: logger.error(message) return UpdateWeightsFromDistributedReqOutput( @@ -160,6 +166,7 @@ class SchedulerWeightUpdaterManager: success, message = worker.update_weights_from_tensor(recv_req) if success: self.flush_cache_after_weight_update(recv_req) + self.record_weight_version_after_update(recv_req.weight_version) else: logger.error(message) torch.distributed.barrier(group=self.tp_cpu_group) @@ -174,7 +181,9 @@ class SchedulerWeightUpdaterManager: success, message = self.draft_worker.update_weights_from_ipc(recv_req) if tp_success: self.flush_cache_after_weight_update(recv_req) - if not success: + if success: + self.record_weight_version_after_update(recv_req.weight_version) + else: logger.error(message) torch.distributed.barrier(group=self.tp_cpu_group) return UpdateWeightsFromIPCReqOutput(success=success, message=message) diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index eed4859e8..994c4b614 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -72,6 +72,8 @@ from sglang.srt.managers.io_struct import ( UpdateWeightsFromIPCReqOutput, UpdateWeightsFromTensorReqInput, UpdateWeightsFromTensorReqOutput, + UpdateWeightVersionReqInput, + UpdateWeightVersionReqOutput, ) from sglang.srt.managers.load_snapshot import LoadSnapshot from sglang.srt.runtime_context import ( @@ -106,6 +108,7 @@ _COMMUNICATOR_SPECS = [ ("send_weights_to_remote_instance", SendWeightsToRemoteInstanceReqOutput), ("update_weights_from_tensor", UpdateWeightsFromTensorReqOutput), ("update_weights_from_ipc", UpdateWeightsFromIPCReqOutput), + ("update_weight_version", UpdateWeightVersionReqOutput), ("get_weights_by_name", GetWeightsByNameReqOutput), ("release_memory_occupation", ReleaseMemoryOccupationReqOutput), ("resume_memory_occupation", ResumeMemoryOccupationReqOutput), @@ -929,6 +932,13 @@ class TokenizerControlMixin: ): await self._async_dispatch_to_scheduler(obj) + async def update_weight_version( + self: TokenizerManager, obj: UpdateWeightVersionReqInput + ) -> None: + self.auto_create_handle_loop() + await self.update_weight_version_communicator(obj) + self._update_weight_version_if_provided(obj.new_version) + def _update_weight_version_if_provided( self: TokenizerManager, weight_version: Optional[str] ) -> None: diff --git a/test/registered/unit/managers/test_scheduler_weight_version_tracking.py b/test/registered/unit/managers/test_scheduler_weight_version_tracking.py new file mode 100644 index 000000000..e304b1fc8 --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_weight_version_tracking.py @@ -0,0 +1,187 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.managers.scheduler_components.weight_updater import ( + SchedulerWeightUpdaterManager, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +class _ServingStub: + def __init__(self, weight_version: str): + self.weight_version = weight_version + + +class _ContextStub: + def __init__(self, serving: _ServingStub): + self.serving = serving + + def override(self, source, **fields): + self.serving.weight_version = fields["weight_version"] + + +class TestSchedulerRecordWeightVersionChange(CustomTestCase): + def _serving(self, version: str) -> _ServingStub: + serving = _ServingStub(version) + for name, value in ( + ("get_serving", serving), + ("get_context", _ContextStub(serving)), + ): + patcher = patch(f"sglang.srt.managers.scheduler.{name}", return_value=value) + patcher.start() + self.addCleanup(patcher.stop) + return serving + + def test_a_new_version_is_adopted(self): + """The scheduler has to end up on the version it was told about, or nothing downstream can read it.""" + serving = self._serving("v1") + + Scheduler.record_weight_version_change(SimpleNamespace(), new_version="v2") + + self.assertEqual(serving.weight_version, "v2") + + def test_same_version_is_a_noop(self): + """Re-announcing the current version must not be treated as a change.""" + serving = self._serving("v1") + + Scheduler.record_weight_version_change(SimpleNamespace(), new_version="v1") + + self.assertEqual(serving.weight_version, "v1") + + def test_none_version_is_a_noop(self): + """An update that carries no version must leave the recorded one alone.""" + serving = self._serving("v1") + + Scheduler.record_weight_version_change(SimpleNamespace(), new_version=None) + + self.assertEqual(serving.weight_version, "v1") + + +class TestRecordWeightVersionAfterUpdate(CustomTestCase): + def _updater( + self, target_result, draft_result=None, method="update_weights_from_disk" + ): + self.recorded = [] + return SchedulerWeightUpdaterManager( + tp_worker=SimpleNamespace(**{method: lambda recv_req: target_result}), + draft_worker=( + None + if draft_result is None + else SimpleNamespace(**{method: lambda recv_req: draft_result}) + ), + tp_cpu_group=None, + memory_saver_adapter=None, + flush_cache=lambda **kwargs: True, + is_fully_idle=lambda **kwargs: True, + scheduler=SimpleNamespace( + record_weight_version_change=lambda new_version: self.recorded.append( + new_version + ) + ), + ) + + def _request(self, **fields): + return SimpleNamespace( + weight_version="v2", + flush_cache=True, + torch_empty_cache=False, + **fields, + ) + + def test_successful_update_records_the_version(self): + """A refit that reports success advances the scheduler-side version.""" + updater = self._updater(target_result=(True, "ok")) + + output = updater.update_weights_from_disk(self._request()) + + self.assertTrue(output.success) + self.assertEqual(self.recorded, ["v2"]) + + def test_failed_update_does_not_record_the_version(self): + """A refit that fails must leave the version alone, or later tokens are mislabelled.""" + updater = self._updater(target_result=(False, "boom")) + + output = updater.update_weights_from_disk(self._request()) + + self.assertFalse(output.success) + self.assertEqual(self.recorded, []) + + def test_draft_failure_does_not_record_the_version(self): + """The target succeeding is not enough: a failed draft refit leaves the engine mixed.""" + updater = self._updater( + target_result=(True, "ok"), draft_result=(False, "draft boom") + ) + + output = updater.update_weights_from_disk(self._request()) + + self.assertFalse(output.success) + self.assertEqual(self.recorded, []) + + def test_successful_distributed_update_records_the_version(self): + """The distributed refit is the path an RL trainer actually drives, so it must record too.""" + updater = self._updater( + target_result=(True, "ok"), method="update_weights_from_distributed" + ) + + output = updater.update_weights_from_distributed(self._request()) + + self.assertTrue(output.success) + self.assertEqual(self.recorded, ["v2"]) + + def test_failed_distributed_update_does_not_record_the_version(self): + """A failed distributed refit leaves the version alone, exactly like the disk path.""" + updater = self._updater( + target_result=(False, "boom"), method="update_weights_from_distributed" + ) + + output = updater.update_weights_from_distributed(self._request()) + + self.assertFalse(output.success) + self.assertEqual(self.recorded, []) + + def test_successful_tensor_update_records_the_version(self): + """The tensor refit records the version once the load reports success.""" + updater = self._updater( + target_result=(True, "ok"), method="update_weights_from_tensor" + ) + + with patch("torch.distributed.barrier"): + output = updater.update_weights_from_tensor( + self._request(disable_draft_model=True) + ) + + self.assertTrue(output.success) + self.assertEqual(self.recorded, ["v2"]) + + def test_successful_ipc_update_records_the_version(self): + """The checkpoint-engine IPC refit records the version like every other path.""" + updater = self._updater( + target_result=(True, "ok"), method="update_weights_from_ipc" + ) + + with patch("torch.distributed.barrier"): + output = updater.update_weights_from_ipc(self._request()) + + self.assertTrue(output.success) + self.assertEqual(self.recorded, ["v2"]) + + def test_failed_ipc_update_does_not_record_the_version(self): + """The IPC path branches on success separately from the cache flush, so failure must record nothing.""" + updater = self._updater( + target_result=(False, "boom"), method="update_weights_from_ipc" + ) + + with patch("torch.distributed.barrier"): + output = updater.update_weights_from_ipc(self._request()) + + self.assertFalse(output.success) + self.assertEqual(self.recorded, []) + + +if __name__ == "__main__": + unittest.main()