Move weight-update RPC handlers to SchedulerWeightUpdaterManager (#25616)
This commit is contained in:
@@ -186,9 +186,6 @@ from sglang.srt.managers.scheduler_runtime_checker_mixin import (
|
|||||||
SchedulerRuntimeCheckerMixin,
|
SchedulerRuntimeCheckerMixin,
|
||||||
create_scheduler_watchdog,
|
create_scheduler_watchdog,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.scheduler_update_weights_mixin import (
|
|
||||||
SchedulerUpdateWeightsMixin,
|
|
||||||
)
|
|
||||||
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
|
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
|
||||||
from sglang.srt.mem_cache import kv_cache_builder
|
from sglang.srt.mem_cache import kv_cache_builder
|
||||||
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
|
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
|
||||||
@@ -324,7 +321,6 @@ def validate_dflash_request(req: Req) -> Optional[str]:
|
|||||||
|
|
||||||
class Scheduler(
|
class Scheduler(
|
||||||
SchedulerOutputProcessorMixin,
|
SchedulerOutputProcessorMixin,
|
||||||
SchedulerUpdateWeightsMixin,
|
|
||||||
SchedulerMetricsMixin,
|
SchedulerMetricsMixin,
|
||||||
SchedulerDisaggregationDecodeMixin,
|
SchedulerDisaggregationDecodeMixin,
|
||||||
SchedulerDisaggregationPrefillMixin,
|
SchedulerDisaggregationPrefillMixin,
|
||||||
@@ -1312,19 +1308,15 @@ class Scheduler(
|
|||||||
(CloseSessionReqInput, self.close_session),
|
(CloseSessionReqInput, self.close_session),
|
||||||
(
|
(
|
||||||
UpdateWeightFromDiskReqInput,
|
UpdateWeightFromDiskReqInput,
|
||||||
lambda req: self.update_weights_from_disk(self.weight_updater, req),
|
self.weight_updater.update_weights_from_disk,
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
InitWeightsUpdateGroupReqInput,
|
InitWeightsUpdateGroupReqInput,
|
||||||
lambda req: self.init_weights_update_group(
|
self.weight_updater.init_weights_update_group,
|
||||||
self.weight_updater, req
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
DestroyWeightsUpdateGroupReqInput,
|
DestroyWeightsUpdateGroupReqInput,
|
||||||
lambda req: self.destroy_weights_update_group(
|
self.weight_updater.destroy_weights_update_group,
|
||||||
self.weight_updater, req
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
InitWeightsSendGroupForRemoteInstanceReqInput,
|
InitWeightsSendGroupForRemoteInstanceReqInput,
|
||||||
@@ -1336,37 +1328,31 @@ class Scheduler(
|
|||||||
),
|
),
|
||||||
(
|
(
|
||||||
UpdateWeightsFromDistributedReqInput,
|
UpdateWeightsFromDistributedReqInput,
|
||||||
lambda req: self.update_weights_from_distributed(
|
self.weight_updater.update_weights_from_distributed,
|
||||||
self.weight_updater, req
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
UpdateWeightsFromTensorReqInput,
|
UpdateWeightsFromTensorReqInput,
|
||||||
lambda req: self.update_weights_from_tensor(
|
self.weight_updater.update_weights_from_tensor,
|
||||||
self.weight_updater, req
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
UpdateWeightsFromIPCReqInput,
|
UpdateWeightsFromIPCReqInput,
|
||||||
lambda req: self.update_weights_from_ipc(self.weight_updater, req),
|
self.weight_updater.update_weights_from_ipc,
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
GetWeightsByNameReqInput,
|
GetWeightsByNameReqInput,
|
||||||
lambda req: self.get_weights_by_name(self.weight_updater, req),
|
self.weight_updater.get_weights_by_name,
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
ReleaseMemoryOccupationReqInput,
|
ReleaseMemoryOccupationReqInput,
|
||||||
lambda req: self.release_memory_occupation(
|
self.weight_updater.release_memory_occupation,
|
||||||
self.weight_updater, req
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
ResumeMemoryOccupationReqInput,
|
ResumeMemoryOccupationReqInput,
|
||||||
lambda req: self.resume_memory_occupation(self.weight_updater, req),
|
self.weight_updater.resume_memory_occupation,
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
CheckWeightsReqInput,
|
CheckWeightsReqInput,
|
||||||
lambda req: self.check_weights(self.weight_updater, req),
|
self.weight_updater.check_weights,
|
||||||
),
|
),
|
||||||
(SlowDownReqInput, self.slow_down),
|
(SlowDownReqInput, self.slow_down),
|
||||||
(
|
(
|
||||||
@@ -3289,10 +3275,10 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
|
|
||||||
def save_remote_model(self, **kwargs):
|
def save_remote_model(self, **kwargs):
|
||||||
SchedulerUpdateWeightsMixin.save_remote_model(self.weight_updater, kwargs)
|
self.weight_updater.save_remote_model(kwargs)
|
||||||
|
|
||||||
def save_sharded_model(self, **kwargs):
|
def save_sharded_model(self, **kwargs):
|
||||||
SchedulerUpdateWeightsMixin.save_sharded_model(self.weight_updater, kwargs)
|
self.weight_updater.save_sharded_model(kwargs)
|
||||||
|
|
||||||
def handle_rpc_request(self, recv_req: RpcReqInput):
|
def handle_rpc_request(self, recv_req: RpcReqInput):
|
||||||
# Handle RPC requests
|
# Handle RPC requests
|
||||||
|
|||||||
@@ -1,7 +1,42 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import traceback
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Callable
|
from typing import Any, Callable, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.constants import (
|
||||||
|
GPU_MEMORY_ALL_TYPES,
|
||||||
|
GPU_MEMORY_TYPE_CUDA_GRAPH,
|
||||||
|
GPU_MEMORY_TYPE_KV_CACHE,
|
||||||
|
GPU_MEMORY_TYPE_WEIGHTS,
|
||||||
|
)
|
||||||
|
from sglang.srt.managers.io_struct import (
|
||||||
|
CheckWeightsReqInput,
|
||||||
|
CheckWeightsReqOutput,
|
||||||
|
DestroyWeightsUpdateGroupReqInput,
|
||||||
|
DestroyWeightsUpdateGroupReqOutput,
|
||||||
|
GetWeightsByNameReqInput,
|
||||||
|
GetWeightsByNameReqOutput,
|
||||||
|
InitWeightsUpdateGroupReqInput,
|
||||||
|
InitWeightsUpdateGroupReqOutput,
|
||||||
|
ReleaseMemoryOccupationReqInput,
|
||||||
|
ReleaseMemoryOccupationReqOutput,
|
||||||
|
ResumeMemoryOccupationReqInput,
|
||||||
|
ResumeMemoryOccupationReqOutput,
|
||||||
|
UpdateWeightFromDiskReqInput,
|
||||||
|
UpdateWeightFromDiskReqOutput,
|
||||||
|
UpdateWeightsFromDistributedReqInput,
|
||||||
|
UpdateWeightsFromDistributedReqOutput,
|
||||||
|
UpdateWeightsFromIPCReqInput,
|
||||||
|
UpdateWeightsFromIPCReqOutput,
|
||||||
|
UpdateWeightsFromTensorReqInput,
|
||||||
|
UpdateWeightsFromTensorReqOutput,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(kw_only=True, slots=True)
|
@dataclass(kw_only=True, slots=True)
|
||||||
@@ -14,3 +49,180 @@ class SchedulerWeightUpdaterManager:
|
|||||||
is_fully_idle: Callable[..., bool]
|
is_fully_idle: Callable[..., bool]
|
||||||
offload_tags: set = field(default_factory=set)
|
offload_tags: set = field(default_factory=set)
|
||||||
stashed_model_static_state: Any = None
|
stashed_model_static_state: Any = None
|
||||||
|
|
||||||
|
def flush_cache_after_weight_update(self, recv_req) -> None:
|
||||||
|
if recv_req.flush_cache:
|
||||||
|
flush_cache_success = self.flush_cache(
|
||||||
|
empty_cache=recv_req.torch_empty_cache
|
||||||
|
)
|
||||||
|
assert flush_cache_success, "Cache flush failed after updating weights"
|
||||||
|
|
||||||
|
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||||
|
"""In-place update of the weights from disk."""
|
||||||
|
success, message = self.tp_worker.update_weights_from_disk(recv_req)
|
||||||
|
tp_success = success
|
||||||
|
if success and self.draft_worker is not None:
|
||||||
|
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:
|
||||||
|
logger.error(message)
|
||||||
|
return UpdateWeightFromDiskReqOutput(success, message, 0)
|
||||||
|
|
||||||
|
def init_weights_update_group(self, recv_req: InitWeightsUpdateGroupReqInput):
|
||||||
|
"""Initialize the online model parameter update group."""
|
||||||
|
success, message = self.tp_worker.init_weights_update_group(recv_req)
|
||||||
|
return InitWeightsUpdateGroupReqOutput(success, message)
|
||||||
|
|
||||||
|
def destroy_weights_update_group(
|
||||||
|
self,
|
||||||
|
recv_req: DestroyWeightsUpdateGroupReqInput,
|
||||||
|
):
|
||||||
|
"""Destroy the online model parameter update group."""
|
||||||
|
success, message = self.tp_worker.destroy_weights_update_group(recv_req)
|
||||||
|
return DestroyWeightsUpdateGroupReqOutput(success, message)
|
||||||
|
|
||||||
|
def update_weights_from_distributed(
|
||||||
|
self,
|
||||||
|
recv_req: UpdateWeightsFromDistributedReqInput,
|
||||||
|
) -> Tuple[bool, str]:
|
||||||
|
"""Update the online model parameter."""
|
||||||
|
success, message = self.tp_worker.update_weights_from_distributed(recv_req)
|
||||||
|
if success:
|
||||||
|
self.flush_cache_after_weight_update(recv_req)
|
||||||
|
else:
|
||||||
|
logger.error(message)
|
||||||
|
return UpdateWeightsFromDistributedReqOutput(success, message)
|
||||||
|
|
||||||
|
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||||
|
"""Update the online model parameter from tensors."""
|
||||||
|
if recv_req.disable_draft_model:
|
||||||
|
worker = self.tp_worker
|
||||||
|
else:
|
||||||
|
worker = self.draft_worker or self.tp_worker
|
||||||
|
success, message = worker.update_weights_from_tensor(recv_req)
|
||||||
|
if success:
|
||||||
|
self.flush_cache_after_weight_update(recv_req)
|
||||||
|
else:
|
||||||
|
logger.error(message)
|
||||||
|
torch.distributed.barrier(group=self.tp_cpu_group)
|
||||||
|
return UpdateWeightsFromTensorReqOutput(success, message)
|
||||||
|
|
||||||
|
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
||||||
|
"""Update the online model parameter from IPC for checkpoint-engine integration."""
|
||||||
|
success, message = self.tp_worker.update_weights_from_ipc(recv_req)
|
||||||
|
tp_success = success
|
||||||
|
if success and self.draft_worker is not None:
|
||||||
|
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:
|
||||||
|
logger.error(message)
|
||||||
|
torch.distributed.barrier(group=self.tp_cpu_group)
|
||||||
|
return UpdateWeightsFromIPCReqOutput(success, message)
|
||||||
|
|
||||||
|
def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput):
|
||||||
|
parameter = self.tp_worker.get_weights_by_name(recv_req)
|
||||||
|
return GetWeightsByNameReqOutput(parameter)
|
||||||
|
|
||||||
|
def release_memory_occupation(self, recv_req: ReleaseMemoryOccupationReqInput):
|
||||||
|
assert (
|
||||||
|
self.is_fully_idle()
|
||||||
|
), "release_memory_occupation should be called only when server is idle."
|
||||||
|
|
||||||
|
tags = recv_req.tags
|
||||||
|
|
||||||
|
if tags is None or len(tags) == 0:
|
||||||
|
tags = GPU_MEMORY_ALL_TYPES
|
||||||
|
|
||||||
|
for tag in tags:
|
||||||
|
self.offload_tags.add(tag)
|
||||||
|
|
||||||
|
if GPU_MEMORY_TYPE_KV_CACHE in tags:
|
||||||
|
self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE)
|
||||||
|
self.flush_cache()
|
||||||
|
|
||||||
|
if GPU_MEMORY_TYPE_WEIGHTS in tags:
|
||||||
|
self.stashed_model_static_state = _export_static_state(
|
||||||
|
self.tp_worker.model_runner.model
|
||||||
|
)
|
||||||
|
torch.distributed.barrier(self.tp_cpu_group)
|
||||||
|
self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_WEIGHTS)
|
||||||
|
|
||||||
|
if GPU_MEMORY_TYPE_CUDA_GRAPH in tags:
|
||||||
|
self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_CUDA_GRAPH)
|
||||||
|
|
||||||
|
torch.get_device_module().synchronize()
|
||||||
|
|
||||||
|
return ReleaseMemoryOccupationReqOutput()
|
||||||
|
|
||||||
|
def resume_memory_occupation(self, recv_req: ResumeMemoryOccupationReqInput):
|
||||||
|
tags = recv_req.tags
|
||||||
|
|
||||||
|
if tags is None or len(tags) == 0:
|
||||||
|
tags = GPU_MEMORY_ALL_TYPES
|
||||||
|
|
||||||
|
for tag in tags:
|
||||||
|
self.offload_tags.remove(tag)
|
||||||
|
|
||||||
|
if GPU_MEMORY_TYPE_CUDA_GRAPH in tags:
|
||||||
|
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_CUDA_GRAPH)
|
||||||
|
|
||||||
|
if GPU_MEMORY_TYPE_WEIGHTS in tags:
|
||||||
|
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_WEIGHTS)
|
||||||
|
torch.distributed.barrier(self.tp_cpu_group)
|
||||||
|
_import_static_state(
|
||||||
|
self.tp_worker.model_runner.model,
|
||||||
|
self.stashed_model_static_state,
|
||||||
|
)
|
||||||
|
del self.stashed_model_static_state
|
||||||
|
|
||||||
|
if GPU_MEMORY_TYPE_KV_CACHE in tags:
|
||||||
|
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE)
|
||||||
|
|
||||||
|
return ResumeMemoryOccupationReqOutput()
|
||||||
|
|
||||||
|
def check_weights(self, recv_req: CheckWeightsReqInput):
|
||||||
|
try:
|
||||||
|
payload = self.tp_worker.model_runner.check_weights(action=recv_req.action)
|
||||||
|
return CheckWeightsReqOutput(
|
||||||
|
success=True, message="Success.", payload=payload
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"check_weights see error: {e}")
|
||||||
|
traceback.print_exc()
|
||||||
|
return CheckWeightsReqOutput(success=False, message=f"{e}")
|
||||||
|
|
||||||
|
def save_remote_model(self, params):
|
||||||
|
url = params["url"]
|
||||||
|
|
||||||
|
self.tp_worker.model_runner.save_remote_model(url)
|
||||||
|
|
||||||
|
if self.draft_worker is not None:
|
||||||
|
draft_url = params.get("draft_url", None)
|
||||||
|
assert (
|
||||||
|
draft_url is not None
|
||||||
|
), "draft_url must be provided when draft model is enabled"
|
||||||
|
self.draft_worker.model_runner.save_remote_model(draft_url)
|
||||||
|
|
||||||
|
def save_sharded_model(self, params):
|
||||||
|
self.tp_worker.model_runner.save_sharded_model(
|
||||||
|
path=params["path"],
|
||||||
|
pattern=params["pattern"],
|
||||||
|
max_size=params["max_size"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _export_static_state(model):
|
||||||
|
return dict(
|
||||||
|
buffers=[
|
||||||
|
(name, buffer.detach().clone()) for name, buffer in model.named_buffers()
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _import_static_state(model, static_params):
|
||||||
|
with torch.inference_mode():
|
||||||
|
self_named_buffers = dict(model.named_buffers())
|
||||||
|
for name, tensor in static_params["buffers"]:
|
||||||
|
self_named_buffers[name][...] = tensor
|
||||||
|
|||||||
@@ -1,253 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
import traceback
|
|
||||||
from typing import TYPE_CHECKING, Tuple
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from sglang.srt.constants import (
|
|
||||||
GPU_MEMORY_ALL_TYPES,
|
|
||||||
GPU_MEMORY_TYPE_CUDA_GRAPH,
|
|
||||||
GPU_MEMORY_TYPE_KV_CACHE,
|
|
||||||
GPU_MEMORY_TYPE_WEIGHTS,
|
|
||||||
)
|
|
||||||
from sglang.srt.managers.io_struct import (
|
|
||||||
CheckWeightsReqInput,
|
|
||||||
CheckWeightsReqOutput,
|
|
||||||
DestroyWeightsUpdateGroupReqInput,
|
|
||||||
DestroyWeightsUpdateGroupReqOutput,
|
|
||||||
GetWeightsByNameReqInput,
|
|
||||||
GetWeightsByNameReqOutput,
|
|
||||||
InitWeightsUpdateGroupReqInput,
|
|
||||||
InitWeightsUpdateGroupReqOutput,
|
|
||||||
ReleaseMemoryOccupationReqInput,
|
|
||||||
ReleaseMemoryOccupationReqOutput,
|
|
||||||
ResumeMemoryOccupationReqInput,
|
|
||||||
ResumeMemoryOccupationReqOutput,
|
|
||||||
UpdateWeightFromDiskReqInput,
|
|
||||||
UpdateWeightFromDiskReqOutput,
|
|
||||||
UpdateWeightsFromDistributedReqInput,
|
|
||||||
UpdateWeightsFromDistributedReqOutput,
|
|
||||||
UpdateWeightsFromIPCReqInput,
|
|
||||||
UpdateWeightsFromIPCReqOutput,
|
|
||||||
UpdateWeightsFromTensorReqInput,
|
|
||||||
UpdateWeightsFromTensorReqOutput,
|
|
||||||
)
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from sglang.srt.managers.scheduler_components.weight_updater import (
|
|
||||||
SchedulerWeightUpdaterManager,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class SchedulerUpdateWeightsMixin:
|
|
||||||
@staticmethod
|
|
||||||
def flush_cache_after_weight_update(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req
|
|
||||||
) -> None:
|
|
||||||
if recv_req.flush_cache:
|
|
||||||
flush_cache_success = self.flush_cache(
|
|
||||||
empty_cache=recv_req.torch_empty_cache
|
|
||||||
)
|
|
||||||
assert flush_cache_success, "Cache flush failed after updating weights"
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def update_weights_from_disk(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: UpdateWeightFromDiskReqInput
|
|
||||||
):
|
|
||||||
"""In-place update of the weights from disk."""
|
|
||||||
success, message = self.tp_worker.update_weights_from_disk(recv_req)
|
|
||||||
tp_success = success
|
|
||||||
if success and self.draft_worker is not None:
|
|
||||||
success, message = self.draft_worker.update_weights_from_disk(recv_req)
|
|
||||||
if tp_success:
|
|
||||||
SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
|
|
||||||
if not success:
|
|
||||||
logger.error(message)
|
|
||||||
return UpdateWeightFromDiskReqOutput(success, message, 0)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def init_weights_update_group(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: InitWeightsUpdateGroupReqInput
|
|
||||||
):
|
|
||||||
"""Initialize the online model parameter update group."""
|
|
||||||
success, message = self.tp_worker.init_weights_update_group(recv_req)
|
|
||||||
return InitWeightsUpdateGroupReqOutput(success, message)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def destroy_weights_update_group(
|
|
||||||
self: "SchedulerWeightUpdaterManager",
|
|
||||||
recv_req: DestroyWeightsUpdateGroupReqInput,
|
|
||||||
):
|
|
||||||
"""Destroy the online model parameter update group."""
|
|
||||||
success, message = self.tp_worker.destroy_weights_update_group(recv_req)
|
|
||||||
return DestroyWeightsUpdateGroupReqOutput(success, message)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def update_weights_from_distributed(
|
|
||||||
self: "SchedulerWeightUpdaterManager",
|
|
||||||
recv_req: UpdateWeightsFromDistributedReqInput,
|
|
||||||
) -> Tuple[bool, str]:
|
|
||||||
"""Update the online model parameter."""
|
|
||||||
success, message = self.tp_worker.update_weights_from_distributed(recv_req)
|
|
||||||
if success:
|
|
||||||
SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
|
|
||||||
else:
|
|
||||||
logger.error(message)
|
|
||||||
return UpdateWeightsFromDistributedReqOutput(success, message)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def update_weights_from_tensor(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: UpdateWeightsFromTensorReqInput
|
|
||||||
):
|
|
||||||
"""Update the online model parameter from tensors."""
|
|
||||||
if recv_req.disable_draft_model:
|
|
||||||
worker = self.tp_worker
|
|
||||||
else:
|
|
||||||
worker = self.draft_worker or self.tp_worker
|
|
||||||
success, message = worker.update_weights_from_tensor(recv_req)
|
|
||||||
if success:
|
|
||||||
SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
|
|
||||||
else:
|
|
||||||
logger.error(message)
|
|
||||||
torch.distributed.barrier(group=self.tp_cpu_group)
|
|
||||||
return UpdateWeightsFromTensorReqOutput(success, message)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def update_weights_from_ipc(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: UpdateWeightsFromIPCReqInput
|
|
||||||
):
|
|
||||||
"""Update the online model parameter from IPC for checkpoint-engine integration."""
|
|
||||||
success, message = self.tp_worker.update_weights_from_ipc(recv_req)
|
|
||||||
tp_success = success
|
|
||||||
if success and self.draft_worker is not None:
|
|
||||||
success, message = self.draft_worker.update_weights_from_ipc(recv_req)
|
|
||||||
if tp_success:
|
|
||||||
SchedulerUpdateWeightsMixin.flush_cache_after_weight_update(self, recv_req)
|
|
||||||
if not success:
|
|
||||||
logger.error(message)
|
|
||||||
torch.distributed.barrier(group=self.tp_cpu_group)
|
|
||||||
return UpdateWeightsFromIPCReqOutput(success, message)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def get_weights_by_name(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: GetWeightsByNameReqInput
|
|
||||||
):
|
|
||||||
parameter = self.tp_worker.get_weights_by_name(recv_req)
|
|
||||||
return GetWeightsByNameReqOutput(parameter)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def release_memory_occupation(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: ReleaseMemoryOccupationReqInput
|
|
||||||
):
|
|
||||||
assert (
|
|
||||||
self.is_fully_idle()
|
|
||||||
), "release_memory_occupation should be called only when server is idle."
|
|
||||||
|
|
||||||
tags = recv_req.tags
|
|
||||||
|
|
||||||
if tags is None or len(tags) == 0:
|
|
||||||
tags = GPU_MEMORY_ALL_TYPES
|
|
||||||
|
|
||||||
for tag in tags:
|
|
||||||
self.offload_tags.add(tag)
|
|
||||||
|
|
||||||
if GPU_MEMORY_TYPE_KV_CACHE in tags:
|
|
||||||
self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE)
|
|
||||||
self.flush_cache()
|
|
||||||
|
|
||||||
if GPU_MEMORY_TYPE_WEIGHTS in tags:
|
|
||||||
self.stashed_model_static_state = _export_static_state(
|
|
||||||
self.tp_worker.model_runner.model
|
|
||||||
)
|
|
||||||
torch.distributed.barrier(self.tp_cpu_group)
|
|
||||||
self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_WEIGHTS)
|
|
||||||
|
|
||||||
if GPU_MEMORY_TYPE_CUDA_GRAPH in tags:
|
|
||||||
self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_CUDA_GRAPH)
|
|
||||||
|
|
||||||
torch.get_device_module().synchronize()
|
|
||||||
|
|
||||||
return ReleaseMemoryOccupationReqOutput()
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def resume_memory_occupation(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: ResumeMemoryOccupationReqInput
|
|
||||||
):
|
|
||||||
tags = recv_req.tags
|
|
||||||
|
|
||||||
if tags is None or len(tags) == 0:
|
|
||||||
tags = GPU_MEMORY_ALL_TYPES
|
|
||||||
|
|
||||||
for tag in tags:
|
|
||||||
self.offload_tags.remove(tag)
|
|
||||||
|
|
||||||
if GPU_MEMORY_TYPE_CUDA_GRAPH in tags:
|
|
||||||
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_CUDA_GRAPH)
|
|
||||||
|
|
||||||
if GPU_MEMORY_TYPE_WEIGHTS in tags:
|
|
||||||
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_WEIGHTS)
|
|
||||||
torch.distributed.barrier(self.tp_cpu_group)
|
|
||||||
_import_static_state(
|
|
||||||
self.tp_worker.model_runner.model,
|
|
||||||
self.stashed_model_static_state,
|
|
||||||
)
|
|
||||||
del self.stashed_model_static_state
|
|
||||||
|
|
||||||
if GPU_MEMORY_TYPE_KV_CACHE in tags:
|
|
||||||
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE)
|
|
||||||
|
|
||||||
return ResumeMemoryOccupationReqOutput()
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def check_weights(
|
|
||||||
self: "SchedulerWeightUpdaterManager", recv_req: CheckWeightsReqInput
|
|
||||||
):
|
|
||||||
try:
|
|
||||||
payload = self.tp_worker.model_runner.check_weights(action=recv_req.action)
|
|
||||||
return CheckWeightsReqOutput(
|
|
||||||
success=True, message="Success.", payload=payload
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"check_weights see error: {e}")
|
|
||||||
traceback.print_exc()
|
|
||||||
return CheckWeightsReqOutput(success=False, message=f"{e}")
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def save_remote_model(self: "SchedulerWeightUpdaterManager", params):
|
|
||||||
url = params["url"]
|
|
||||||
|
|
||||||
self.tp_worker.model_runner.save_remote_model(url)
|
|
||||||
|
|
||||||
if self.draft_worker is not None:
|
|
||||||
draft_url = params.get("draft_url", None)
|
|
||||||
assert (
|
|
||||||
draft_url is not None
|
|
||||||
), "draft_url must be provided when draft model is enabled"
|
|
||||||
self.draft_worker.model_runner.save_remote_model(draft_url)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def save_sharded_model(self: "SchedulerWeightUpdaterManager", params):
|
|
||||||
self.tp_worker.model_runner.save_sharded_model(
|
|
||||||
path=params["path"],
|
|
||||||
pattern=params["pattern"],
|
|
||||||
max_size=params["max_size"],
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _export_static_state(model):
|
|
||||||
return dict(
|
|
||||||
buffers=[
|
|
||||||
(name, buffer.detach().clone()) for name, buffer in model.named_buffers()
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _import_static_state(model, static_params):
|
|
||||||
with torch.inference_mode():
|
|
||||||
self_named_buffers = dict(model.named_buffers())
|
|
||||||
for name, tensor in static_params["buffers"]:
|
|
||||||
self_named_buffers[name][...] = tensor
|
|
||||||
Reference in New Issue
Block a user