Carve out SchedulerWeightUpdaterManager for weight-update state (#25615)

This commit is contained in:
fzyzcjy
2026-05-18 18:33:48 +08:00
committed by GitHub
parent a35690f070
commit 56f27635b8
3 changed files with 122 additions and 30 deletions
+66 -12
View File
@@ -173,6 +173,9 @@ from sglang.srt.managers.scheduler_components.profiler_manager import (
from sglang.srt.managers.scheduler_components.request_receiver import ( from sglang.srt.managers.scheduler_components.request_receiver import (
SchedulerRequestReceiver, SchedulerRequestReceiver,
) )
from sglang.srt.managers.scheduler_components.weight_updater import (
SchedulerWeightUpdaterManager,
)
from sglang.srt.managers.scheduler_input_blocker import SchedulerInputBlocker from sglang.srt.managers.scheduler_input_blocker import SchedulerInputBlocker
from sglang.srt.managers.scheduler_output_processor_mixin import ( from sglang.srt.managers.scheduler_output_processor_mixin import (
SchedulerOutputProcessorMixin, SchedulerOutputProcessorMixin,
@@ -545,6 +548,15 @@ class Scheduler(
# Init prefill kv split size when deterministic inference is enabled with various attention backends # Init prefill kv split size when deterministic inference is enabled with various attention backends
self.init_deterministic_inference_config() self.init_deterministic_inference_config()
self.weight_updater = SchedulerWeightUpdaterManager(
tp_worker=self.tp_worker,
draft_worker=self.draft_worker,
tp_cpu_group=self.tp_cpu_group,
memory_saver_adapter=self.memory_saver_adapter,
flush_cache=self.flush_cache,
is_fully_idle=self.is_fully_idle,
)
# Init request dispatcher # Init request dispatcher
self.init_request_dispatcher() self.init_request_dispatcher()
@@ -594,7 +606,7 @@ class Scheduler(
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
offload_tags=self.offload_tags, offload_tags=self.weight_updater.offload_tags,
ps=self.ps, ps=self.ps,
server_args=self.server_args, server_args=self.server_args,
model_config=self.model_config, model_config=self.model_config,
@@ -1039,7 +1051,6 @@ class Scheduler(
self.memory_saver_adapter = TorchMemorySaverAdapter.create( self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=self.server_args.enable_memory_saver enable=self.server_args.enable_memory_saver
) )
self.offload_tags = set()
# Init recv skipper and input blocker # Init recv skipper and input blocker
self.recv_skipper = SchedulerRecvSkipper.maybe_create(self.server_args) self.recv_skipper = SchedulerRecvSkipper.maybe_create(self.server_args)
@@ -1299,9 +1310,22 @@ class Scheduler(
(AbortReq, self.abort_request), (AbortReq, self.abort_request),
(OpenSessionReqInput, self.open_session), (OpenSessionReqInput, self.open_session),
(CloseSessionReqInput, self.close_session), (CloseSessionReqInput, self.close_session),
(UpdateWeightFromDiskReqInput, self.update_weights_from_disk), (
(InitWeightsUpdateGroupReqInput, self.init_weights_update_group), UpdateWeightFromDiskReqInput,
(DestroyWeightsUpdateGroupReqInput, self.destroy_weights_update_group), lambda req: self.update_weights_from_disk(self.weight_updater, req),
),
(
InitWeightsUpdateGroupReqInput,
lambda req: self.init_weights_update_group(
self.weight_updater, req
),
),
(
DestroyWeightsUpdateGroupReqInput,
lambda req: self.destroy_weights_update_group(
self.weight_updater, req
),
),
( (
InitWeightsSendGroupForRemoteInstanceReqInput, InitWeightsSendGroupForRemoteInstanceReqInput,
self.init_weights_send_group_for_remote_instance, self.init_weights_send_group_for_remote_instance,
@@ -1312,14 +1336,38 @@ class Scheduler(
), ),
( (
UpdateWeightsFromDistributedReqInput, UpdateWeightsFromDistributedReqInput,
self.update_weights_from_distributed, lambda req: self.update_weights_from_distributed(
self.weight_updater, req
),
),
(
UpdateWeightsFromTensorReqInput,
lambda req: self.update_weights_from_tensor(
self.weight_updater, req
),
),
(
UpdateWeightsFromIPCReqInput,
lambda req: self.update_weights_from_ipc(self.weight_updater, req),
),
(
GetWeightsByNameReqInput,
lambda req: self.get_weights_by_name(self.weight_updater, req),
),
(
ReleaseMemoryOccupationReqInput,
lambda req: self.release_memory_occupation(
self.weight_updater, req
),
),
(
ResumeMemoryOccupationReqInput,
lambda req: self.resume_memory_occupation(self.weight_updater, req),
),
(
CheckWeightsReqInput,
lambda req: self.check_weights(self.weight_updater, req),
), ),
(UpdateWeightsFromTensorReqInput, self.update_weights_from_tensor),
(UpdateWeightsFromIPCReqInput, self.update_weights_from_ipc),
(GetWeightsByNameReqInput, self.get_weights_by_name),
(ReleaseMemoryOccupationReqInput, self.release_memory_occupation),
(ResumeMemoryOccupationReqInput, self.resume_memory_occupation),
(CheckWeightsReqInput, self.check_weights),
(SlowDownReqInput, self.slow_down), (SlowDownReqInput, self.slow_down),
( (
ProfileReq, ProfileReq,
@@ -3240,6 +3288,12 @@ class Scheduler(
server_args=vars(get_global_server_args()), server_args=vars(get_global_server_args()),
) )
def save_remote_model(self, **kwargs):
SchedulerUpdateWeightsMixin.save_remote_model(self.weight_updater, kwargs)
def save_sharded_model(self, **kwargs):
SchedulerUpdateWeightsMixin.save_sharded_model(self.weight_updater, kwargs)
def handle_rpc_request(self, recv_req: RpcReqInput): def handle_rpc_request(self, recv_req: RpcReqInput):
# Handle RPC requests # Handle RPC requests
logger.info( logger.info(
@@ -0,0 +1,16 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Callable
@dataclass(kw_only=True, slots=True)
class SchedulerWeightUpdaterManager:
tp_worker: Any
draft_worker: Any
tp_cpu_group: Any
memory_saver_adapter: Any
flush_cache: Callable[..., bool]
is_fully_idle: Callable[..., bool]
offload_tags: set = field(default_factory=set)
stashed_model_static_state: Any = None
@@ -36,21 +36,27 @@ from sglang.srt.managers.io_struct import (
) )
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler_components.weight_updater import (
SchedulerWeightUpdaterManager,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class SchedulerUpdateWeightsMixin: class SchedulerUpdateWeightsMixin:
def flush_cache_after_weight_update(self: Scheduler, recv_req) -> None: @staticmethod
def flush_cache_after_weight_update(
self: "SchedulerWeightUpdaterManager", recv_req
) -> None:
if recv_req.flush_cache: if recv_req.flush_cache:
flush_cache_success = self.flush_cache( flush_cache_success = self.flush_cache(
empty_cache=recv_req.torch_empty_cache empty_cache=recv_req.torch_empty_cache
) )
assert flush_cache_success, "Cache flush failed after updating weights" assert flush_cache_success, "Cache flush failed after updating weights"
@staticmethod
def update_weights_from_disk( def update_weights_from_disk(
self: Scheduler, recv_req: UpdateWeightFromDiskReqInput self: "SchedulerWeightUpdaterManager", recv_req: UpdateWeightFromDiskReqInput
): ):
"""In-place update of the weights from disk.""" """In-place update of the weights from disk."""
success, message = self.tp_worker.update_weights_from_disk(recv_req) success, message = self.tp_worker.update_weights_from_disk(recv_req)
@@ -58,39 +64,44 @@ class SchedulerUpdateWeightsMixin:
if success and self.draft_worker is not None: if success and self.draft_worker is not None:
success, message = self.draft_worker.update_weights_from_disk(recv_req) success, message = self.draft_worker.update_weights_from_disk(recv_req)
if tp_success: if tp_success:
self.flush_cache_after_weight_update(recv_req) SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
if not success: if not success:
logger.error(message) logger.error(message)
return UpdateWeightFromDiskReqOutput(success, message, 0) return UpdateWeightFromDiskReqOutput(success, message, 0)
@staticmethod
def init_weights_update_group( def init_weights_update_group(
self: Scheduler, recv_req: InitWeightsUpdateGroupReqInput self: "SchedulerWeightUpdaterManager", recv_req: InitWeightsUpdateGroupReqInput
): ):
"""Initialize the online model parameter update group.""" """Initialize the online model parameter update group."""
success, message = self.tp_worker.init_weights_update_group(recv_req) success, message = self.tp_worker.init_weights_update_group(recv_req)
return InitWeightsUpdateGroupReqOutput(success, message) return InitWeightsUpdateGroupReqOutput(success, message)
@staticmethod
def destroy_weights_update_group( def destroy_weights_update_group(
self: Scheduler, recv_req: DestroyWeightsUpdateGroupReqInput self: "SchedulerWeightUpdaterManager",
recv_req: DestroyWeightsUpdateGroupReqInput,
): ):
"""Destroy the online model parameter update group.""" """Destroy the online model parameter update group."""
success, message = self.tp_worker.destroy_weights_update_group(recv_req) success, message = self.tp_worker.destroy_weights_update_group(recv_req)
return DestroyWeightsUpdateGroupReqOutput(success, message) return DestroyWeightsUpdateGroupReqOutput(success, message)
@staticmethod
def update_weights_from_distributed( def update_weights_from_distributed(
self, self: "SchedulerWeightUpdaterManager",
recv_req: UpdateWeightsFromDistributedReqInput, recv_req: UpdateWeightsFromDistributedReqInput,
) -> Tuple[bool, str]: ) -> Tuple[bool, str]:
"""Update the online model parameter.""" """Update the online model parameter."""
success, message = self.tp_worker.update_weights_from_distributed(recv_req) success, message = self.tp_worker.update_weights_from_distributed(recv_req)
if success: if success:
self.flush_cache_after_weight_update(recv_req) SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
else: else:
logger.error(message) logger.error(message)
return UpdateWeightsFromDistributedReqOutput(success, message) return UpdateWeightsFromDistributedReqOutput(success, message)
@staticmethod
def update_weights_from_tensor( def update_weights_from_tensor(
self: Scheduler, recv_req: UpdateWeightsFromTensorReqInput self: "SchedulerWeightUpdaterManager", recv_req: UpdateWeightsFromTensorReqInput
): ):
"""Update the online model parameter from tensors.""" """Update the online model parameter from tensors."""
if recv_req.disable_draft_model: if recv_req.disable_draft_model:
@@ -99,14 +110,15 @@ class SchedulerUpdateWeightsMixin:
worker = self.draft_worker or self.tp_worker worker = self.draft_worker or self.tp_worker
success, message = worker.update_weights_from_tensor(recv_req) success, message = worker.update_weights_from_tensor(recv_req)
if success: if success:
self.flush_cache_after_weight_update(recv_req) SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
else: else:
logger.error(message) logger.error(message)
torch.distributed.barrier(group=self.tp_cpu_group) torch.distributed.barrier(group=self.tp_cpu_group)
return UpdateWeightsFromTensorReqOutput(success, message) return UpdateWeightsFromTensorReqOutput(success, message)
@staticmethod
def update_weights_from_ipc( def update_weights_from_ipc(
self: Scheduler, recv_req: UpdateWeightsFromIPCReqInput self: "SchedulerWeightUpdaterManager", recv_req: UpdateWeightsFromIPCReqInput
): ):
"""Update the online model parameter from IPC for checkpoint-engine integration.""" """Update the online model parameter from IPC for checkpoint-engine integration."""
success, message = self.tp_worker.update_weights_from_ipc(recv_req) success, message = self.tp_worker.update_weights_from_ipc(recv_req)
@@ -114,18 +126,22 @@ class SchedulerUpdateWeightsMixin:
if success and self.draft_worker is not None: if success and self.draft_worker is not None:
success, message = self.draft_worker.update_weights_from_ipc(recv_req) success, message = self.draft_worker.update_weights_from_ipc(recv_req)
if tp_success: if tp_success:
self.flush_cache_after_weight_update(recv_req) SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
if not success: if not success:
logger.error(message) logger.error(message)
torch.distributed.barrier(group=self.tp_cpu_group) torch.distributed.barrier(group=self.tp_cpu_group)
return UpdateWeightsFromIPCReqOutput(success, message) return UpdateWeightsFromIPCReqOutput(success, message)
def get_weights_by_name(self: Scheduler, recv_req: GetWeightsByNameReqInput): @staticmethod
def get_weights_by_name(
self: "SchedulerWeightUpdaterManager", recv_req: GetWeightsByNameReqInput
):
parameter = self.tp_worker.get_weights_by_name(recv_req) parameter = self.tp_worker.get_weights_by_name(recv_req)
return GetWeightsByNameReqOutput(parameter) return GetWeightsByNameReqOutput(parameter)
@staticmethod
def release_memory_occupation( def release_memory_occupation(
self: Scheduler, recv_req: ReleaseMemoryOccupationReqInput self: "SchedulerWeightUpdaterManager", recv_req: ReleaseMemoryOccupationReqInput
): ):
assert ( assert (
self.is_fully_idle() self.is_fully_idle()
@@ -157,8 +173,9 @@ class SchedulerUpdateWeightsMixin:
return ReleaseMemoryOccupationReqOutput() return ReleaseMemoryOccupationReqOutput()
@staticmethod
def resume_memory_occupation( def resume_memory_occupation(
self: Scheduler, recv_req: ResumeMemoryOccupationReqInput self: "SchedulerWeightUpdaterManager", recv_req: ResumeMemoryOccupationReqInput
): ):
tags = recv_req.tags tags = recv_req.tags
@@ -185,7 +202,10 @@ class SchedulerUpdateWeightsMixin:
return ResumeMemoryOccupationReqOutput() return ResumeMemoryOccupationReqOutput()
def check_weights(self: Scheduler, recv_req: CheckWeightsReqInput): @staticmethod
def check_weights(
self: "SchedulerWeightUpdaterManager", recv_req: CheckWeightsReqInput
):
try: try:
payload = self.tp_worker.model_runner.check_weights(action=recv_req.action) payload = self.tp_worker.model_runner.check_weights(action=recv_req.action)
return CheckWeightsReqOutput( return CheckWeightsReqOutput(
@@ -196,7 +216,8 @@ class SchedulerUpdateWeightsMixin:
traceback.print_exc() traceback.print_exc()
return CheckWeightsReqOutput(success=False, message=f"{e}") return CheckWeightsReqOutput(success=False, message=f"{e}")
def save_remote_model(self: Scheduler, params): @staticmethod
def save_remote_model(self: "SchedulerWeightUpdaterManager", params):
url = params["url"] url = params["url"]
self.tp_worker.model_runner.save_remote_model(url) self.tp_worker.model_runner.save_remote_model(url)
@@ -208,7 +229,8 @@ class SchedulerUpdateWeightsMixin:
), "draft_url must be provided when draft model is enabled" ), "draft_url must be provided when draft model is enabled"
self.draft_worker.model_runner.save_remote_model(draft_url) self.draft_worker.model_runner.save_remote_model(draft_url)
def save_sharded_model(self: Scheduler, params): @staticmethod
def save_sharded_model(self: "SchedulerWeightUpdaterManager", params):
self.tp_worker.model_runner.save_sharded_model( self.tp_worker.model_runner.save_sharded_model(
path=params["path"], path=params["path"],
pattern=params["pattern"], pattern=params["pattern"],