Make the scheduler track the published weight version (#35925)
This commit is contained in:
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user