Carve out SchedulerWeightUpdaterManager for weight-update state (#25615)
This commit is contained in:
@@ -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"],
|
||||||
|
|||||||
Reference in New Issue
Block a user