Introduce WeightUpdater and WeightExporter components (#31148)
This commit is contained in:
@@ -300,17 +300,17 @@ class SchedulerWeightUpdaterManager:
|
|||||||
def save_remote_model(self, params):
|
def save_remote_model(self, params):
|
||||||
url = params["url"]
|
url = params["url"]
|
||||||
|
|
||||||
self.tp_worker.model_runner.save_remote_model(url)
|
self.tp_worker.model_runner.weight_exporter.save_remote_model(url)
|
||||||
|
|
||||||
if self.draft_worker is not None:
|
if self.draft_worker is not None:
|
||||||
draft_url = params.get("draft_url", None)
|
draft_url = params.get("draft_url", None)
|
||||||
assert (
|
assert (
|
||||||
draft_url is not None
|
draft_url is not None
|
||||||
), "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.weight_exporter.save_remote_model(draft_url)
|
||||||
|
|
||||||
def save_sharded_model(self, params):
|
def save_sharded_model(self, params):
|
||||||
self.tp_worker.model_runner.save_sharded_model(
|
self.tp_worker.model_runner.weight_exporter.save_sharded_model(
|
||||||
path=params["path"],
|
path=params["path"],
|
||||||
pattern=params["pattern"],
|
pattern=params["pattern"],
|
||||||
max_size=params["max_size"],
|
max_size=params["max_size"],
|
||||||
|
|||||||
@@ -101,7 +101,7 @@ class BaseTpWorker(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||||
success, message = self.model_runner.update_weights_from_disk(
|
success, message = self.model_runner.weight_updater.update_weights_from_disk(
|
||||||
recv_req.model_path,
|
recv_req.model_path,
|
||||||
recv_req.load_format,
|
recv_req.load_format,
|
||||||
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
||||||
@@ -109,7 +109,7 @@ class BaseTpWorker(ABC):
|
|||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
def init_weights_update_group(self, recv_req: InitWeightsUpdateGroupReqInput):
|
def init_weights_update_group(self, recv_req: InitWeightsUpdateGroupReqInput):
|
||||||
success, message = self.model_runner.init_weights_update_group(
|
success, message = self.model_runner.weight_updater.init_weights_update_group(
|
||||||
recv_req.master_address,
|
recv_req.master_address,
|
||||||
recv_req.master_port,
|
recv_req.master_port,
|
||||||
recv_req.rank_offset,
|
recv_req.rank_offset,
|
||||||
@@ -120,8 +120,10 @@ class BaseTpWorker(ABC):
|
|||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
def destroy_weights_update_group(self, recv_req: DestroyWeightsUpdateGroupReqInput):
|
def destroy_weights_update_group(self, recv_req: DestroyWeightsUpdateGroupReqInput):
|
||||||
success, message = self.model_runner.destroy_weights_update_group(
|
success, message = (
|
||||||
recv_req.group_name,
|
self.model_runner.weight_updater.destroy_weights_update_group(
|
||||||
|
recv_req.group_name,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
@@ -129,7 +131,7 @@ class BaseTpWorker(ABC):
|
|||||||
self, recv_req: InitWeightsSendGroupForRemoteInstanceReqInput
|
self, recv_req: InitWeightsSendGroupForRemoteInstanceReqInput
|
||||||
):
|
):
|
||||||
success, message = (
|
success, message = (
|
||||||
self.model_runner.init_weights_send_group_for_remote_instance(
|
self.model_runner.weight_exporter.init_weights_send_group_for_remote_instance(
|
||||||
recv_req.master_address,
|
recv_req.master_address,
|
||||||
recv_req.ports,
|
recv_req.ports,
|
||||||
recv_req.group_rank,
|
recv_req.group_rank,
|
||||||
@@ -143,29 +145,33 @@ class BaseTpWorker(ABC):
|
|||||||
def send_weights_to_remote_instance(
|
def send_weights_to_remote_instance(
|
||||||
self, recv_req: SendWeightsToRemoteInstanceReqInput
|
self, recv_req: SendWeightsToRemoteInstanceReqInput
|
||||||
):
|
):
|
||||||
success, message = self.model_runner.send_weights_to_remote_instance(
|
success, message = (
|
||||||
recv_req.master_address,
|
self.model_runner.weight_exporter.send_weights_to_remote_instance(
|
||||||
recv_req.ports,
|
recv_req.master_address,
|
||||||
recv_req.group_name,
|
recv_req.ports,
|
||||||
|
recv_req.group_name,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
def update_weights_from_distributed(
|
def update_weights_from_distributed(
|
||||||
self, recv_req: UpdateWeightsFromDistributedReqInput
|
self, recv_req: UpdateWeightsFromDistributedReqInput
|
||||||
):
|
):
|
||||||
success, message = self.model_runner.update_weights_from_distributed(
|
success, message = (
|
||||||
recv_req.names,
|
self.model_runner.weight_updater.update_weights_from_distributed(
|
||||||
recv_req.dtypes,
|
recv_req.names,
|
||||||
recv_req.shapes,
|
recv_req.dtypes,
|
||||||
recv_req.group_name,
|
recv_req.shapes,
|
||||||
recv_req.load_format,
|
recv_req.group_name,
|
||||||
|
recv_req.load_format,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||||
|
|
||||||
monkey_patch_torch_reductions()
|
monkey_patch_torch_reductions()
|
||||||
success, message = self.model_runner.update_weights_from_tensor(
|
success, message = self.model_runner.weight_updater.update_weights_from_tensor(
|
||||||
named_tensors=MultiprocessingSerializer.deserialize(
|
named_tensors=MultiprocessingSerializer.deserialize(
|
||||||
recv_req.serialized_named_tensors[self.tp_rank]
|
recv_req.serialized_named_tensors[self.tp_rank]
|
||||||
),
|
),
|
||||||
@@ -175,11 +181,13 @@ class BaseTpWorker(ABC):
|
|||||||
|
|
||||||
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
||||||
"""Update weights from IPC for checkpoint-engine integration."""
|
"""Update weights from IPC for checkpoint-engine integration."""
|
||||||
success, message = self.model_runner.update_weights_from_ipc(recv_req)
|
success, message = self.model_runner.weight_updater.update_weights_from_ipc(
|
||||||
|
recv_req
|
||||||
|
)
|
||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput):
|
def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput):
|
||||||
parameter = self.model_runner.get_weights_by_name(
|
parameter = self.model_runner.weight_exporter.get_weights_by_name(
|
||||||
recv_req.name, recv_req.truncate_size
|
recv_req.name, recv_req.truncate_size
|
||||||
)
|
)
|
||||||
return parameter
|
return parameter
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import contextlib
|
import contextlib
|
||||||
import datetime
|
import datetime
|
||||||
import gc
|
|
||||||
import inspect
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -26,7 +25,7 @@ import threading
|
|||||||
import time
|
import time
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Callable, List, Optional, Tuple, Union
|
from typing import Any, List, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
@@ -159,6 +158,12 @@ from sglang.srt.model_executor.forward_context import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
|
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
|
||||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
||||||
|
from sglang.srt.model_executor.model_runner_components.weight_exporter import (
|
||||||
|
WeightExporter,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.model_runner_components.weight_updater import (
|
||||||
|
WeightUpdater,
|
||||||
|
)
|
||||||
from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
|
from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
|
||||||
ModelRunnerKVCacheMixin,
|
ModelRunnerKVCacheMixin,
|
||||||
)
|
)
|
||||||
@@ -171,17 +176,13 @@ from sglang.srt.model_executor.runner import (
|
|||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
|
from sglang.srt.model_loader.loader import get_model_loader
|
||||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||||
RemoteInstanceWeightLoaderBackend,
|
RemoteInstanceWeightLoaderBackend,
|
||||||
register_memory_region,
|
register_memory_region,
|
||||||
trigger_init_weights_send_group_for_remote_instance_request,
|
trigger_init_weights_send_group_for_remote_instance_request,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_loader.utils import (
|
from sglang.srt.model_loader.utils import resolve_language_model
|
||||||
resolve_language_model,
|
|
||||||
set_default_torch_dtype,
|
|
||||||
)
|
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args
|
||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
@@ -206,15 +207,12 @@ from sglang.srt.state_capturer.routed_experts import (
|
|||||||
set_global_experts_capturer,
|
set_global_experts_capturer,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
MultiprocessingSerializer,
|
|
||||||
broadcast_pyobj,
|
broadcast_pyobj,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
dynamic_import,
|
|
||||||
enable_show_time_cost,
|
enable_show_time_cost,
|
||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
init_cublas,
|
init_cublas,
|
||||||
init_custom_process_group,
|
|
||||||
is_hip,
|
is_hip,
|
||||||
is_host_cpu_arm64,
|
is_host_cpu_arm64,
|
||||||
is_npu,
|
is_npu,
|
||||||
@@ -234,17 +232,10 @@ from sglang.srt.utils.offloader import (
|
|||||||
get_offloader,
|
get_offloader,
|
||||||
set_offloader,
|
set_offloader,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.patch_torch import (
|
from sglang.srt.utils.patch_torch import register_sgl_tp_rank
|
||||||
monkey_patch_torch_reductions,
|
|
||||||
register_sgl_tp_rank,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils.profile_utils import build_step_span_name
|
from sglang.srt.utils.profile_utils import build_step_span_name
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
from sglang.srt.utils.weight_checker import WeightChecker
|
from sglang.srt.utils.weight_checker import WeightChecker
|
||||||
from sglang.srt.weight_sync.tensor_bucket import (
|
|
||||||
FlattenedTensorBucket,
|
|
||||||
FlattenedTensorMetadata,
|
|
||||||
)
|
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
@@ -577,8 +568,30 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
), "Pipeline Parallel is not compatible with this model."
|
), "Pipeline Parallel is not compatible with this model."
|
||||||
|
|
||||||
# For weight updates
|
# For weight updates
|
||||||
self._model_update_group = {}
|
self.init_weight_updater()
|
||||||
self._weights_send_group = {}
|
self.init_weight_exporter()
|
||||||
|
|
||||||
|
def init_weight_updater(self):
|
||||||
|
self.weight_updater = WeightUpdater(
|
||||||
|
tp_rank=self.tp_rank,
|
||||||
|
device=self.device,
|
||||||
|
gpu_id=self.gpu_id,
|
||||||
|
model_config=self.model_config,
|
||||||
|
custom_weight_loaders=self.server_args.custom_weight_loader,
|
||||||
|
get_model=lambda: self.model,
|
||||||
|
update_model_fields=self.update_model_fields,
|
||||||
|
recapture_cuda_graph=self.init_decode_cuda_graph,
|
||||||
|
get_model_runner=lambda: self,
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_weight_exporter(self):
|
||||||
|
self.weight_exporter = WeightExporter(
|
||||||
|
tp_rank=self.tp_rank,
|
||||||
|
tp_size=self.tp_size,
|
||||||
|
gpu_id=self.gpu_id,
|
||||||
|
get_model_path=lambda: self.model_config.model_path,
|
||||||
|
get_model=lambda: self.model,
|
||||||
|
)
|
||||||
|
|
||||||
def init_msprobe(self):
|
def init_msprobe(self):
|
||||||
# Init the msprobe
|
# Init the msprobe
|
||||||
@@ -1672,7 +1685,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.expert_backup_client.update_weights(weight_name_filter)
|
self.expert_backup_client.update_weights(weight_name_filter)
|
||||||
else:
|
else:
|
||||||
# Load the missing weights from disk
|
# Load the missing weights from disk
|
||||||
self.update_weights_from_disk(
|
self.weight_updater.update_weights_from_disk(
|
||||||
get_server_args().model_path,
|
get_server_args().model_path,
|
||||||
get_server_args().load_format,
|
get_server_args().load_format,
|
||||||
weight_name_filter=weight_name_filter,
|
weight_name_filter=weight_name_filter,
|
||||||
@@ -1742,409 +1755,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
"All ranks are marked as elastic_ep_rejoin."
|
"All ranks are marked as elastic_ep_rejoin."
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_weights_from_disk(
|
|
||||||
self,
|
|
||||||
model_path: str,
|
|
||||||
load_format: str,
|
|
||||||
weight_name_filter: Optional[Callable[[str], bool]] = None,
|
|
||||||
recapture_cuda_graph: bool = False,
|
|
||||||
) -> tuple[bool, str]:
|
|
||||||
"""Update engine weights in-place from the disk."""
|
|
||||||
logger.info(
|
|
||||||
f"Update engine weights online from disk begin. "
|
|
||||||
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id, empty_cache=False):.2f} GB"
|
|
||||||
)
|
|
||||||
|
|
||||||
target_device = torch.device(self.device)
|
|
||||||
self.model_config.model_path = model_path
|
|
||||||
load_config = LoadConfig(load_format=load_format)
|
|
||||||
|
|
||||||
# Only support DefaultModelLoader for now
|
|
||||||
loader = get_model_loader(load_config, self.model_config)
|
|
||||||
if not isinstance(loader, DefaultModelLoader):
|
|
||||||
message = f"Failed to get model loader: {loader}."
|
|
||||||
return False, message
|
|
||||||
|
|
||||||
def get_weight_iter(config):
|
|
||||||
iter = loader._get_weights_iterator(
|
|
||||||
DefaultModelLoader.Source.init_new(config, self.model)
|
|
||||||
)
|
|
||||||
if weight_name_filter is not None:
|
|
||||||
iter = (
|
|
||||||
(name, weight) for name, weight in iter if weight_name_filter(name)
|
|
||||||
)
|
|
||||||
|
|
||||||
return iter
|
|
||||||
|
|
||||||
def model_load_weights(model, iter):
|
|
||||||
loader.load_weights_and_postprocess(model, iter, target_device)
|
|
||||||
return model
|
|
||||||
|
|
||||||
with set_default_torch_dtype(self.model_config.dtype):
|
|
||||||
try:
|
|
||||||
iter = get_weight_iter(self.model_config)
|
|
||||||
except Exception as e:
|
|
||||||
message = f"Failed to get weights iterator: {e}."
|
|
||||||
return False, message
|
|
||||||
try:
|
|
||||||
model = model_load_weights(self.model, iter)
|
|
||||||
except Exception as e:
|
|
||||||
message = (
|
|
||||||
f"Failed to update weights: {e}.\nRolling back to original weights."
|
|
||||||
)
|
|
||||||
del iter
|
|
||||||
gc.collect()
|
|
||||||
iter = get_weight_iter(self.model_config)
|
|
||||||
self.model = model_load_weights(self.model, iter)
|
|
||||||
return False, message
|
|
||||||
|
|
||||||
self.model = model
|
|
||||||
self.server_args.override(
|
|
||||||
"model_runner.update_weights",
|
|
||||||
model_path=model_path,
|
|
||||||
load_format=load_format,
|
|
||||||
)
|
|
||||||
self.load_config = load_config
|
|
||||||
|
|
||||||
if recapture_cuda_graph and (
|
|
||||||
self.device == "cuda"
|
|
||||||
or self.device == "musa"
|
|
||||||
or (
|
|
||||||
current_platform.is_out_of_tree()
|
|
||||||
and current_platform.support_cuda_graph()
|
|
||||||
)
|
|
||||||
):
|
|
||||||
self.init_decode_cuda_graph()
|
|
||||||
|
|
||||||
logger.info("Update weights end.")
|
|
||||||
return True, "Succeeded to update model weights."
|
|
||||||
|
|
||||||
def init_weights_send_group_for_remote_instance(
|
|
||||||
self,
|
|
||||||
master_address,
|
|
||||||
ports,
|
|
||||||
group_rank,
|
|
||||||
world_size,
|
|
||||||
group_name,
|
|
||||||
backend="nccl",
|
|
||||||
):
|
|
||||||
assert (
|
|
||||||
torch.distributed.is_initialized()
|
|
||||||
), "Default torch process group must be initialized"
|
|
||||||
assert group_name != "", "Group name cannot be empty"
|
|
||||||
|
|
||||||
ports_list = ports.split(",")
|
|
||||||
assert (
|
|
||||||
len(ports_list) == self.tp_size
|
|
||||||
), f"Expected {self.tp_size} ports, but got {len(ports_list)} ports."
|
|
||||||
group_port = ports_list[self.tp_rank]
|
|
||||||
group_name = f"{group_name}_{group_port}_{self.tp_rank}"
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
f"init custom process group: tp_rank={self.tp_rank}, gpu_id={self.gpu_id}, master_address={master_address}, master_port={group_port}, "
|
|
||||||
f"group_rank={group_rank}, world_size={world_size}, group_name={group_name}, backend={backend}"
|
|
||||||
)
|
|
||||||
|
|
||||||
current_platform.empty_cache()
|
|
||||||
success = False
|
|
||||||
message = ""
|
|
||||||
try:
|
|
||||||
na = NetworkAddress(master_address, group_port)
|
|
||||||
self._weights_send_group[group_name] = init_custom_process_group(
|
|
||||||
backend=backend,
|
|
||||||
init_method=na.to_tcp(),
|
|
||||||
world_size=world_size,
|
|
||||||
rank=group_rank,
|
|
||||||
group_name=group_name,
|
|
||||||
device_id=torch.device("cuda", self.gpu_id),
|
|
||||||
)
|
|
||||||
dist.barrier(group=self._weights_send_group[group_name])
|
|
||||||
success = True
|
|
||||||
message = f"Succeeded to init group through {na.to_host_port_str()} group."
|
|
||||||
except Exception as e:
|
|
||||||
message = f"Failed to init group: {e}."
|
|
||||||
logger.error(message)
|
|
||||||
|
|
||||||
current_platform.empty_cache()
|
|
||||||
return success, message
|
|
||||||
|
|
||||||
def send_weights_to_remote_instance(
|
|
||||||
self,
|
|
||||||
master_address,
|
|
||||||
ports,
|
|
||||||
group_name,
|
|
||||||
):
|
|
||||||
assert (
|
|
||||||
torch.distributed.is_initialized()
|
|
||||||
), "Default torch process group must be initialized"
|
|
||||||
assert group_name != "", "Group name cannot be empty"
|
|
||||||
|
|
||||||
ports_list = ports.split(",")
|
|
||||||
assert (
|
|
||||||
len(ports_list) == self.tp_size
|
|
||||||
), f"Expected {self.tp_size} ports, but got {len(ports_list)} ports."
|
|
||||||
group_port = ports_list[self.tp_rank]
|
|
||||||
group_name = f"{group_name}_{group_port}_{self.tp_rank}"
|
|
||||||
|
|
||||||
if self._weights_send_group[group_name] is not None:
|
|
||||||
send_group = self._weights_send_group[group_name]
|
|
||||||
else:
|
|
||||||
message = f"Group {group_name} not in _weights_send_group list. Please call `init_weights_send_group_for_remote_instance` first."
|
|
||||||
logger.error(message)
|
|
||||||
return False, message
|
|
||||||
|
|
||||||
current_platform.empty_cache()
|
|
||||||
success = False
|
|
||||||
na = NetworkAddress(master_address, group_port)
|
|
||||||
message = ""
|
|
||||||
try:
|
|
||||||
for _, weights in self.model.named_parameters():
|
|
||||||
torch.distributed.broadcast(
|
|
||||||
weights,
|
|
||||||
src=0,
|
|
||||||
group=send_group,
|
|
||||||
)
|
|
||||||
success = True
|
|
||||||
message = f"Succeeded to send weights through {na.to_host_port_str()} {group_name}."
|
|
||||||
except Exception as e:
|
|
||||||
message = f"Failed to send weights: {e}."
|
|
||||||
logger.error(message)
|
|
||||||
|
|
||||||
# destroy the process group after sending weights
|
|
||||||
del self._weights_send_group[group_name]
|
|
||||||
torch.distributed.distributed_c10d.destroy_process_group(send_group)
|
|
||||||
current_platform.empty_cache()
|
|
||||||
return success, message
|
|
||||||
|
|
||||||
def init_weights_update_group(
|
|
||||||
self,
|
|
||||||
master_address,
|
|
||||||
master_port,
|
|
||||||
rank_offset,
|
|
||||||
world_size,
|
|
||||||
group_name,
|
|
||||||
backend="nccl",
|
|
||||||
):
|
|
||||||
"""Initialize the Torch process group for model parameter updates.
|
|
||||||
|
|
||||||
`_model_update_group` is used in the RLHF workflow, where rank
|
|
||||||
0 is the actor model in the training engine, and the other ranks are
|
|
||||||
the inference engine, which is used for rollout.
|
|
||||||
|
|
||||||
In the RLHF workflow, the training engine updates the model
|
|
||||||
weights/parameters online, and broadcasts them to the inference
|
|
||||||
engine through the `_model_update_group` process group.
|
|
||||||
"""
|
|
||||||
assert (
|
|
||||||
torch.distributed.is_initialized()
|
|
||||||
), "Default torch process group must be initialized"
|
|
||||||
assert group_name != "", "Group name cannot be empty"
|
|
||||||
|
|
||||||
rank = rank_offset + self.tp_rank
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
f"init custom process group: master_address={master_address}, master_port={master_port}, "
|
|
||||||
f"rank_offset={rank_offset}, rank={rank}, world_size={world_size}, group_name={group_name}, backend={backend}"
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
na = NetworkAddress(master_address, master_port)
|
|
||||||
self._model_update_group[group_name] = init_custom_process_group(
|
|
||||||
backend=backend,
|
|
||||||
init_method=na.to_tcp(),
|
|
||||||
world_size=world_size,
|
|
||||||
rank=rank,
|
|
||||||
group_name=group_name,
|
|
||||||
)
|
|
||||||
return True, "Succeeded to initialize custom process group."
|
|
||||||
except Exception as e:
|
|
||||||
message = f"Failed to initialize custom process group: {e}."
|
|
||||||
logger.error(message)
|
|
||||||
return False, message
|
|
||||||
|
|
||||||
def destroy_weights_update_group(self, group_name):
|
|
||||||
try:
|
|
||||||
if group_name in self._model_update_group:
|
|
||||||
pg = self._model_update_group.pop(group_name)
|
|
||||||
torch.distributed.destroy_process_group(pg)
|
|
||||||
return True, "Succeeded to destroy custom process group."
|
|
||||||
else:
|
|
||||||
return False, "The group to be destroyed does not exist."
|
|
||||||
except Exception as e:
|
|
||||||
message = f"Failed to destroy custom process group: {e}."
|
|
||||||
logger.error(message)
|
|
||||||
return False, message
|
|
||||||
|
|
||||||
def update_weights_from_distributed(
|
|
||||||
self,
|
|
||||||
names,
|
|
||||||
dtypes,
|
|
||||||
shapes,
|
|
||||||
group_name,
|
|
||||||
load_format: Optional[str] = None,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Update specific parameter in the model weights online
|
|
||||||
through `_model_update_group` process group.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
name: the name of the parameter to be updated.
|
|
||||||
dtype: the data type of the parameter to be updated.
|
|
||||||
shape: the shape of the parameter to be updated.
|
|
||||||
"""
|
|
||||||
|
|
||||||
assert group_name in self._model_update_group, (
|
|
||||||
f"Group {group_name} not in {list(self._model_update_group.keys())}. "
|
|
||||||
"Please call `init_weights_update_group` first."
|
|
||||||
)
|
|
||||||
|
|
||||||
if load_format == "flattened_bucket":
|
|
||||||
return self._update_bucketed_weights_from_distributed(
|
|
||||||
names, dtypes, shapes, group_name
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
weights = []
|
|
||||||
handles = []
|
|
||||||
for name, dtype, shape in zip(names, dtypes, shapes):
|
|
||||||
target_dtype = (
|
|
||||||
dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
|
|
||||||
)
|
|
||||||
weight = torch.empty(shape, dtype=target_dtype, device=self.device)
|
|
||||||
handles.append(
|
|
||||||
torch.distributed.broadcast(
|
|
||||||
weight,
|
|
||||||
src=0,
|
|
||||||
group=self._model_update_group[group_name],
|
|
||||||
async_op=True,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
weights.append((name, weight))
|
|
||||||
for handle in handles:
|
|
||||||
handle.wait()
|
|
||||||
|
|
||||||
self.model.load_weights(weights)
|
|
||||||
return True, "Succeeded to update parameter online."
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
error_msg = (
|
|
||||||
f"Failed to update parameter online: {e}. "
|
|
||||||
f"The full weights of the ModelRunner are partially updated. "
|
|
||||||
f"Please discard the whole weights."
|
|
||||||
)
|
|
||||||
logger.error(error_msg)
|
|
||||||
return False, error_msg
|
|
||||||
|
|
||||||
def _update_bucketed_weights_from_distributed(
|
|
||||||
self, names, dtypes, shapes, group_name
|
|
||||||
):
|
|
||||||
try:
|
|
||||||
named_tensors = []
|
|
||||||
for name, dtype, shape in zip(names, dtypes, shapes):
|
|
||||||
target_dtype = (
|
|
||||||
dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
|
|
||||||
)
|
|
||||||
named_tensors.append(
|
|
||||||
(name, torch.empty(shape, dtype=target_dtype, device=self.device))
|
|
||||||
)
|
|
||||||
bucket = FlattenedTensorBucket(named_tensors=named_tensors)
|
|
||||||
flattened_tensor = bucket.get_flattened_tensor()
|
|
||||||
torch.distributed.broadcast(
|
|
||||||
flattened_tensor,
|
|
||||||
src=0,
|
|
||||||
group=self._model_update_group[group_name],
|
|
||||||
)
|
|
||||||
reconstructed_tensors = bucket.reconstruct_tensors()
|
|
||||||
self.model.load_weights(reconstructed_tensors)
|
|
||||||
return True, f"Succeeded to update parameter online."
|
|
||||||
except Exception as e:
|
|
||||||
error_msg = (
|
|
||||||
f"Failed to update parameter online: {e}. "
|
|
||||||
f"The full weights of the ModelRunner are partially updated. "
|
|
||||||
f"Please discard the whole weights."
|
|
||||||
)
|
|
||||||
logger.error(error_msg)
|
|
||||||
return False, error_msg
|
|
||||||
|
|
||||||
def update_weights_from_tensor(
|
|
||||||
self,
|
|
||||||
named_tensors: List[Tuple[str, Union[torch.Tensor, LocalSerializedTensor]]],
|
|
||||||
load_format: Optional[str] = None,
|
|
||||||
):
|
|
||||||
monkey_patch_torch_reductions()
|
|
||||||
if load_format == "flattened_bucket":
|
|
||||||
# Handle flattened bucket format
|
|
||||||
return self._update_weights_from_flattened_bucket(
|
|
||||||
flattened_tensor_bucket_dict=named_tensors
|
|
||||||
)
|
|
||||||
|
|
||||||
# We need to get device after patch otherwise the device would be wrong
|
|
||||||
device_module = torch.get_device_module(self.device)
|
|
||||||
infered_device = device_module.current_device()
|
|
||||||
|
|
||||||
named_tensors = [
|
|
||||||
(name, _unwrap_tensor(tensor, tp_rank=self.tp_rank, device=infered_device))
|
|
||||||
for name, tensor in named_tensors
|
|
||||||
]
|
|
||||||
if load_format == "direct":
|
|
||||||
_model_load_weights_direct(self.model, named_tensors)
|
|
||||||
elif load_format in self.server_args.custom_weight_loader:
|
|
||||||
custom_loader = dynamic_import(load_format)
|
|
||||||
custom_loader(self.model, named_tensors)
|
|
||||||
elif load_format is None:
|
|
||||||
self.model.load_weights(named_tensors)
|
|
||||||
else:
|
|
||||||
raise NotImplementedError(f"Unknown load_format={load_format}")
|
|
||||||
return True, "Success"
|
|
||||||
|
|
||||||
def _update_weights_from_flattened_bucket(
|
|
||||||
self,
|
|
||||||
flattened_tensor_bucket_dict,
|
|
||||||
):
|
|
||||||
"""Handle flattened bucket format for weight updates"""
|
|
||||||
flattened_tensor = flattened_tensor_bucket_dict["flattened_tensor"]
|
|
||||||
metadata = flattened_tensor_bucket_dict["metadata"]
|
|
||||||
|
|
||||||
# Convert metadata dict to our format
|
|
||||||
converted_metadata = []
|
|
||||||
for meta in metadata:
|
|
||||||
converted_meta = FlattenedTensorMetadata(
|
|
||||||
name=meta.name,
|
|
||||||
shape=meta.shape,
|
|
||||||
dtype=meta.dtype,
|
|
||||||
start_idx=meta.start_idx,
|
|
||||||
end_idx=meta.end_idx,
|
|
||||||
numel=meta.numel,
|
|
||||||
)
|
|
||||||
converted_metadata.append(converted_meta)
|
|
||||||
|
|
||||||
# Create bucket and reconstruct tensors
|
|
||||||
bucket = FlattenedTensorBucket(
|
|
||||||
flattened_tensor=flattened_tensor, metadata=converted_metadata
|
|
||||||
)
|
|
||||||
reconstructed_tensors = bucket.reconstruct_tensors()
|
|
||||||
|
|
||||||
# Load the reconstructed tensors using the standard method
|
|
||||||
self.model.load_weights(reconstructed_tensors)
|
|
||||||
|
|
||||||
return True, "Success"
|
|
||||||
|
|
||||||
def get_weights_by_name(
|
|
||||||
self, name: str, truncate_size: int = 100
|
|
||||||
) -> Optional[torch.Tensor]:
|
|
||||||
"""Get the weights of the parameter by its name. Similar to `get_parameter` in Hugging Face.
|
|
||||||
|
|
||||||
Only used for unit test with an unoptimized performance.
|
|
||||||
For optimized performance, please use torch.save and torch.load.
|
|
||||||
"""
|
|
||||||
# TODO: (chenyang) Add support for Qwen models.
|
|
||||||
try:
|
|
||||||
return self.model.get_weights_by_name(
|
|
||||||
name, truncate_size, tp_size=self.tp_size
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error when getting parameter {name}: {e}")
|
|
||||||
return None
|
|
||||||
|
|
||||||
def init_lora_manager(self):
|
def init_lora_manager(self):
|
||||||
self.lora_manager = LoRAManager(
|
self.lora_manager = LoRAManager(
|
||||||
base_model=self.model,
|
base_model=self.model,
|
||||||
@@ -3176,44 +2786,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
forward_batch.token_ids_logprobs,
|
forward_batch.token_ids_logprobs,
|
||||||
)
|
)
|
||||||
|
|
||||||
def save_remote_model(self, url: str):
|
|
||||||
from sglang.srt.model_loader.loader import RemoteModelLoader
|
|
||||||
|
|
||||||
logger.info(f"Saving model to {url}")
|
|
||||||
RemoteModelLoader.save_model(self.model, self.model_config.model_path, url)
|
|
||||||
|
|
||||||
def save_sharded_model(
|
|
||||||
self, path: str, pattern: Optional[str] = None, max_size: Optional[int] = None
|
|
||||||
):
|
|
||||||
from sglang.srt.model_loader.loader import ShardedStateLoader
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
f"Save sharded model to {path} with pattern {pattern} and max_size {max_size}"
|
|
||||||
)
|
|
||||||
ShardedStateLoader.save_model(self.model, path, pattern, max_size)
|
|
||||||
|
|
||||||
def check_weights(self, action: str, allow_quant_error: bool = False):
|
def check_weights(self, action: str, allow_quant_error: bool = False):
|
||||||
return self._weight_checker.handle(
|
return self._weight_checker.handle(
|
||||||
action=action, allow_quant_error=allow_quant_error
|
action=action, allow_quant_error=allow_quant_error
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_weights_from_ipc(self, recv_req):
|
|
||||||
"""Update weights from IPC for checkpoint-engine integration."""
|
|
||||||
try:
|
|
||||||
from sglang.srt.checkpoint_engine.checkpoint_engine_worker import (
|
|
||||||
SGLangCheckpointEngineWorkerExtensionImpl,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Create a worker extension that integrates with SGLang's model
|
|
||||||
worker = SGLangCheckpointEngineWorkerExtensionImpl(self)
|
|
||||||
worker.update_weights_from_ipc(recv_req.zmq_handles)
|
|
||||||
return True, "IPC weight update completed successfully"
|
|
||||||
except ImportError as e:
|
|
||||||
return False, f"IPC weight update failed: ImportError {e}"
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"IPC weight update failed: {e}")
|
|
||||||
return False, str(e)
|
|
||||||
|
|
||||||
def _maybe_rebalance_after_rank_fault(
|
def _maybe_rebalance_after_rank_fault(
|
||||||
self,
|
self,
|
||||||
output: ModelRunnerOutput,
|
output: ModelRunnerOutput,
|
||||||
@@ -3241,25 +2818,18 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
def update_model_fields(
|
||||||
def _model_load_weights_direct(model, named_tensors: List[Tuple[str, torch.Tensor]]):
|
self,
|
||||||
params_dict = dict(model.named_parameters())
|
new_model: torch.nn.Module,
|
||||||
for name, tensor in named_tensors:
|
*,
|
||||||
default_weight_loader(params_dict[name], tensor)
|
model_path: str,
|
||||||
|
load_format: str,
|
||||||
|
load_config: LoadConfig,
|
||||||
def _unwrap_tensor(tensor, tp_rank, device):
|
) -> None:
|
||||||
if isinstance(tensor, LocalSerializedTensor):
|
self.model = new_model
|
||||||
tensor = tensor.get(tp_rank)
|
self.server_args.override(
|
||||||
return tensor.to(device)
|
"model_runner.update_model_fields",
|
||||||
|
model_path=model_path,
|
||||||
|
load_format=load_format,
|
||||||
@dataclass
|
)
|
||||||
class LocalSerializedTensor:
|
self.load_config = load_config
|
||||||
"""torch.Tensor that gets serialized by MultiprocessingSerializer (which only serializes a pointer and not the data).
|
|
||||||
The i-th element in the list corresponds to i-th rank's GPU."""
|
|
||||||
|
|
||||||
values: List[bytes]
|
|
||||||
|
|
||||||
def get(self, rank: int):
|
|
||||||
return MultiprocessingSerializer.deserialize(self.values[rank])
|
|
||||||
|
|||||||
@@ -0,0 +1,157 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Callable, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
|
from sglang.srt.utils import init_custom_process_group
|
||||||
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True, kw_only=True)
|
||||||
|
class WeightExporter:
|
||||||
|
tp_rank: int
|
||||||
|
tp_size: int
|
||||||
|
gpu_id: int
|
||||||
|
get_model_path: Callable[[], str]
|
||||||
|
get_model: Callable[[], Any]
|
||||||
|
_weights_send_group: dict = field(default_factory=dict)
|
||||||
|
|
||||||
|
def init_weights_send_group_for_remote_instance(
|
||||||
|
self,
|
||||||
|
master_address,
|
||||||
|
ports,
|
||||||
|
group_rank,
|
||||||
|
world_size,
|
||||||
|
group_name,
|
||||||
|
backend="nccl",
|
||||||
|
):
|
||||||
|
assert (
|
||||||
|
torch.distributed.is_initialized()
|
||||||
|
), "Default torch process group must be initialized"
|
||||||
|
assert group_name != "", "Group name cannot be empty"
|
||||||
|
|
||||||
|
ports_list = ports.split(",")
|
||||||
|
assert (
|
||||||
|
len(ports_list) == self.tp_size
|
||||||
|
), f"Expected {self.tp_size} ports, but got {len(ports_list)} ports."
|
||||||
|
group_port = ports_list[self.tp_rank]
|
||||||
|
group_name = f"{group_name}_{group_port}_{self.tp_rank}"
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"init custom process group: tp_rank={self.tp_rank}, gpu_id={self.gpu_id}, master_address={master_address}, master_port={group_port}, "
|
||||||
|
f"group_rank={group_rank}, world_size={world_size}, group_name={group_name}, backend={backend}"
|
||||||
|
)
|
||||||
|
|
||||||
|
current_platform.empty_cache()
|
||||||
|
success = False
|
||||||
|
message = ""
|
||||||
|
try:
|
||||||
|
na = NetworkAddress(master_address, group_port)
|
||||||
|
self._weights_send_group[group_name] = init_custom_process_group(
|
||||||
|
backend=backend,
|
||||||
|
init_method=na.to_tcp(),
|
||||||
|
world_size=world_size,
|
||||||
|
rank=group_rank,
|
||||||
|
group_name=group_name,
|
||||||
|
device_id=torch.device("cuda", self.gpu_id),
|
||||||
|
)
|
||||||
|
dist.barrier(group=self._weights_send_group[group_name])
|
||||||
|
success = True
|
||||||
|
message = f"Succeeded to init group through {na.to_host_port_str()} group."
|
||||||
|
except Exception as e:
|
||||||
|
message = f"Failed to init group: {e}."
|
||||||
|
logger.error(message)
|
||||||
|
|
||||||
|
current_platform.empty_cache()
|
||||||
|
return success, message
|
||||||
|
|
||||||
|
def send_weights_to_remote_instance(
|
||||||
|
self,
|
||||||
|
master_address,
|
||||||
|
ports,
|
||||||
|
group_name,
|
||||||
|
):
|
||||||
|
assert (
|
||||||
|
torch.distributed.is_initialized()
|
||||||
|
), "Default torch process group must be initialized"
|
||||||
|
assert group_name != "", "Group name cannot be empty"
|
||||||
|
|
||||||
|
ports_list = ports.split(",")
|
||||||
|
assert (
|
||||||
|
len(ports_list) == self.tp_size
|
||||||
|
), f"Expected {self.tp_size} ports, but got {len(ports_list)} ports."
|
||||||
|
group_port = ports_list[self.tp_rank]
|
||||||
|
group_name = f"{group_name}_{group_port}_{self.tp_rank}"
|
||||||
|
|
||||||
|
if self._weights_send_group[group_name] is not None:
|
||||||
|
send_group = self._weights_send_group[group_name]
|
||||||
|
else:
|
||||||
|
message = f"Group {group_name} not in _weights_send_group list. Please call `init_weights_send_group_for_remote_instance` first."
|
||||||
|
logger.error(message)
|
||||||
|
return False, message
|
||||||
|
|
||||||
|
current_platform.empty_cache()
|
||||||
|
success = False
|
||||||
|
na = NetworkAddress(master_address, group_port)
|
||||||
|
message = ""
|
||||||
|
try:
|
||||||
|
for _, weights in self.get_model().named_parameters():
|
||||||
|
torch.distributed.broadcast(
|
||||||
|
weights,
|
||||||
|
src=0,
|
||||||
|
group=send_group,
|
||||||
|
)
|
||||||
|
success = True
|
||||||
|
message = f"Succeeded to send weights through {na.to_host_port_str()} {group_name}."
|
||||||
|
except Exception as e:
|
||||||
|
message = f"Failed to send weights: {e}."
|
||||||
|
logger.error(message)
|
||||||
|
|
||||||
|
# destroy the process group after sending weights
|
||||||
|
del self._weights_send_group[group_name]
|
||||||
|
torch.distributed.distributed_c10d.destroy_process_group(send_group)
|
||||||
|
current_platform.empty_cache()
|
||||||
|
return success, message
|
||||||
|
|
||||||
|
def save_remote_model(self: WeightExporter, url: str):
|
||||||
|
from sglang.srt.model_loader.loader import RemoteModelLoader
|
||||||
|
|
||||||
|
logger.info(f"Saving model to {url}")
|
||||||
|
RemoteModelLoader.save_model(self.get_model(), self.get_model_path(), url)
|
||||||
|
|
||||||
|
def save_sharded_model(
|
||||||
|
self: WeightExporter,
|
||||||
|
path: str,
|
||||||
|
pattern: Optional[str] = None,
|
||||||
|
max_size: Optional[int] = None,
|
||||||
|
):
|
||||||
|
from sglang.srt.model_loader.loader import ShardedStateLoader
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Save sharded model to {path} with pattern {pattern} and max_size {max_size}"
|
||||||
|
)
|
||||||
|
ShardedStateLoader.save_model(self.get_model(), path, pattern, max_size)
|
||||||
|
|
||||||
|
def get_weights_by_name(
|
||||||
|
self: WeightExporter, name: str, truncate_size: int = 100
|
||||||
|
) -> Optional[torch.Tensor]:
|
||||||
|
"""Get the weights of the parameter by its name. Similar to `get_parameter` in Hugging Face.
|
||||||
|
|
||||||
|
Only used for unit test with an unoptimized performance.
|
||||||
|
For optimized performance, please use torch.save and torch.load.
|
||||||
|
"""
|
||||||
|
# TODO: (chenyang) Add support for Qwen models.
|
||||||
|
try:
|
||||||
|
return self.get_model().get_weights_by_name(
|
||||||
|
name, truncate_size, tp_size=self.tp_size
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error when getting parameter {name}: {e}")
|
||||||
|
return None
|
||||||
@@ -0,0 +1,377 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import gc
|
||||||
|
import logging
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import TYPE_CHECKING, Any, Callable, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.configs.load_config import LoadConfig
|
||||||
|
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
|
||||||
|
from sglang.srt.model_loader.utils import set_default_torch_dtype
|
||||||
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
|
from sglang.srt.utils import (
|
||||||
|
MultiprocessingSerializer,
|
||||||
|
dynamic_import,
|
||||||
|
get_available_gpu_memory,
|
||||||
|
init_custom_process_group,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
||||||
|
from sglang.srt.weight_sync.tensor_bucket import (
|
||||||
|
FlattenedTensorBucket,
|
||||||
|
FlattenedTensorMetadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||||
|
class WeightUpdater:
|
||||||
|
tp_rank: int
|
||||||
|
device: str
|
||||||
|
gpu_id: int
|
||||||
|
model_config: ModelConfig
|
||||||
|
custom_weight_loaders: dict
|
||||||
|
get_model: Callable[[], Any]
|
||||||
|
update_model_fields: Callable[..., None]
|
||||||
|
recapture_cuda_graph: Callable[[], None]
|
||||||
|
get_model_runner: Callable[[], ModelRunner]
|
||||||
|
_model_update_group: dict = field(default_factory=dict)
|
||||||
|
|
||||||
|
def init_weights_update_group(
|
||||||
|
self,
|
||||||
|
master_address,
|
||||||
|
master_port,
|
||||||
|
rank_offset,
|
||||||
|
world_size,
|
||||||
|
group_name,
|
||||||
|
backend="nccl",
|
||||||
|
):
|
||||||
|
"""Initialize the Torch process group for model parameter updates.
|
||||||
|
|
||||||
|
`_model_update_group` is used in the RLHF workflow, where rank
|
||||||
|
0 is the actor model in the training engine, and the other ranks are
|
||||||
|
the inference engine, which is used for rollout.
|
||||||
|
|
||||||
|
In the RLHF workflow, the training engine updates the model
|
||||||
|
weights/parameters online, and broadcasts them to the inference
|
||||||
|
engine through the `_model_update_group` process group.
|
||||||
|
"""
|
||||||
|
assert (
|
||||||
|
torch.distributed.is_initialized()
|
||||||
|
), "Default torch process group must be initialized"
|
||||||
|
assert group_name != "", "Group name cannot be empty"
|
||||||
|
|
||||||
|
rank = rank_offset + self.tp_rank
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"init custom process group: master_address={master_address}, master_port={master_port}, "
|
||||||
|
f"rank_offset={rank_offset}, rank={rank}, world_size={world_size}, group_name={group_name}, backend={backend}"
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
na = NetworkAddress(master_address, master_port)
|
||||||
|
self._model_update_group[group_name] = init_custom_process_group(
|
||||||
|
backend=backend,
|
||||||
|
init_method=na.to_tcp(),
|
||||||
|
world_size=world_size,
|
||||||
|
rank=rank,
|
||||||
|
group_name=group_name,
|
||||||
|
)
|
||||||
|
return True, "Succeeded to initialize custom process group."
|
||||||
|
except Exception as e:
|
||||||
|
message = f"Failed to initialize custom process group: {e}."
|
||||||
|
logger.error(message)
|
||||||
|
return False, message
|
||||||
|
|
||||||
|
def destroy_weights_update_group(self, group_name):
|
||||||
|
try:
|
||||||
|
if group_name in self._model_update_group:
|
||||||
|
pg = self._model_update_group.pop(group_name)
|
||||||
|
torch.distributed.destroy_process_group(pg)
|
||||||
|
return True, "Succeeded to destroy custom process group."
|
||||||
|
else:
|
||||||
|
return False, "The group to be destroyed does not exist."
|
||||||
|
except Exception as e:
|
||||||
|
message = f"Failed to destroy custom process group: {e}."
|
||||||
|
logger.error(message)
|
||||||
|
return False, message
|
||||||
|
|
||||||
|
def update_weights_from_disk(
|
||||||
|
self: WeightUpdater,
|
||||||
|
model_path: str,
|
||||||
|
load_format: str,
|
||||||
|
weight_name_filter: Optional[Callable[[str], bool]] = None,
|
||||||
|
recapture_cuda_graph: bool = False,
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
"""Update engine weights in-place from the disk."""
|
||||||
|
logger.info(
|
||||||
|
f"Update engine weights online from disk begin. "
|
||||||
|
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id, empty_cache=False):.2f} GB"
|
||||||
|
)
|
||||||
|
|
||||||
|
target_device = torch.device(self.device)
|
||||||
|
self.model_config.model_path = model_path
|
||||||
|
load_config = LoadConfig(load_format=load_format)
|
||||||
|
|
||||||
|
# Only support DefaultModelLoader for now
|
||||||
|
loader = get_model_loader(load_config, self.model_config)
|
||||||
|
if not isinstance(loader, DefaultModelLoader):
|
||||||
|
message = f"Failed to get model loader: {loader}."
|
||||||
|
return False, message
|
||||||
|
|
||||||
|
def get_weight_iter(config):
|
||||||
|
iter = loader._get_weights_iterator(
|
||||||
|
DefaultModelLoader.Source.init_new(config, self.get_model())
|
||||||
|
)
|
||||||
|
if weight_name_filter is not None:
|
||||||
|
iter = (
|
||||||
|
(name, weight) for name, weight in iter if weight_name_filter(name)
|
||||||
|
)
|
||||||
|
|
||||||
|
return iter
|
||||||
|
|
||||||
|
def model_load_weights(model, iter):
|
||||||
|
loader.load_weights_and_postprocess(model, iter, target_device)
|
||||||
|
return model
|
||||||
|
|
||||||
|
with set_default_torch_dtype(self.model_config.dtype):
|
||||||
|
try:
|
||||||
|
iter = get_weight_iter(self.model_config)
|
||||||
|
except Exception as e:
|
||||||
|
message = f"Failed to get weights iterator: {e}."
|
||||||
|
return False, message
|
||||||
|
try:
|
||||||
|
model = model_load_weights(self.get_model(), iter)
|
||||||
|
except Exception as e:
|
||||||
|
message = (
|
||||||
|
f"Failed to update weights: {e}.\nRolling back to original weights."
|
||||||
|
)
|
||||||
|
del iter
|
||||||
|
gc.collect()
|
||||||
|
iter = get_weight_iter(self.model_config)
|
||||||
|
model_load_weights(self.get_model(), iter)
|
||||||
|
return False, message
|
||||||
|
|
||||||
|
self.update_model_fields(
|
||||||
|
model,
|
||||||
|
model_path=model_path,
|
||||||
|
load_format=load_format,
|
||||||
|
load_config=load_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
if recapture_cuda_graph and (
|
||||||
|
self.device == "cuda"
|
||||||
|
or self.device == "musa"
|
||||||
|
or (
|
||||||
|
current_platform.is_out_of_tree()
|
||||||
|
and current_platform.support_cuda_graph()
|
||||||
|
)
|
||||||
|
):
|
||||||
|
self.recapture_cuda_graph()
|
||||||
|
|
||||||
|
logger.info("Update weights end.")
|
||||||
|
return True, "Succeeded to update model weights."
|
||||||
|
|
||||||
|
def update_weights_from_distributed(
|
||||||
|
self: WeightUpdater,
|
||||||
|
names,
|
||||||
|
dtypes,
|
||||||
|
shapes,
|
||||||
|
group_name,
|
||||||
|
load_format: Optional[str] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Update specific parameter in the model weights online
|
||||||
|
through `_model_update_group` process group.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name: the name of the parameter to be updated.
|
||||||
|
dtype: the data type of the parameter to be updated.
|
||||||
|
shape: the shape of the parameter to be updated.
|
||||||
|
"""
|
||||||
|
|
||||||
|
assert group_name in self._model_update_group, (
|
||||||
|
f"Group {group_name} not in {list(self._model_update_group.keys())}. "
|
||||||
|
"Please call `init_weights_update_group` first."
|
||||||
|
)
|
||||||
|
|
||||||
|
if load_format == "flattened_bucket":
|
||||||
|
return self._update_bucketed_weights_from_distributed(
|
||||||
|
names, dtypes, shapes, group_name
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
weights = []
|
||||||
|
handles = []
|
||||||
|
for name, dtype, shape in zip(names, dtypes, shapes):
|
||||||
|
target_dtype = (
|
||||||
|
dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
|
||||||
|
)
|
||||||
|
weight = torch.empty(shape, dtype=target_dtype, device=self.device)
|
||||||
|
handles.append(
|
||||||
|
torch.distributed.broadcast(
|
||||||
|
weight,
|
||||||
|
src=0,
|
||||||
|
group=self._model_update_group[group_name],
|
||||||
|
async_op=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
weights.append((name, weight))
|
||||||
|
for handle in handles:
|
||||||
|
handle.wait()
|
||||||
|
|
||||||
|
self.get_model().load_weights(weights)
|
||||||
|
return True, "Succeeded to update parameter online."
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
error_msg = (
|
||||||
|
f"Failed to update parameter online: {e}. "
|
||||||
|
f"The full weights of the ModelRunner are partially updated. "
|
||||||
|
f"Please discard the whole weights."
|
||||||
|
)
|
||||||
|
logger.error(error_msg)
|
||||||
|
return False, error_msg
|
||||||
|
|
||||||
|
def _update_bucketed_weights_from_distributed(
|
||||||
|
self: WeightUpdater, names, dtypes, shapes, group_name
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
named_tensors = []
|
||||||
|
for name, dtype, shape in zip(names, dtypes, shapes):
|
||||||
|
target_dtype = (
|
||||||
|
dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
|
||||||
|
)
|
||||||
|
named_tensors.append(
|
||||||
|
(
|
||||||
|
name,
|
||||||
|
torch.empty(shape, dtype=target_dtype, device=self.device),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
bucket = FlattenedTensorBucket(named_tensors=named_tensors)
|
||||||
|
flattened_tensor = bucket.get_flattened_tensor()
|
||||||
|
torch.distributed.broadcast(
|
||||||
|
flattened_tensor,
|
||||||
|
src=0,
|
||||||
|
group=self._model_update_group[group_name],
|
||||||
|
)
|
||||||
|
reconstructed_tensors = bucket.reconstruct_tensors()
|
||||||
|
self.get_model().load_weights(reconstructed_tensors)
|
||||||
|
return True, f"Succeeded to update parameter online."
|
||||||
|
except Exception as e:
|
||||||
|
error_msg = (
|
||||||
|
f"Failed to update parameter online: {e}. "
|
||||||
|
f"The full weights of the ModelRunner are partially updated. "
|
||||||
|
f"Please discard the whole weights."
|
||||||
|
)
|
||||||
|
logger.error(error_msg)
|
||||||
|
return False, error_msg
|
||||||
|
|
||||||
|
def update_weights_from_tensor(
|
||||||
|
self: WeightUpdater,
|
||||||
|
named_tensors: List[Tuple[str, Union[torch.Tensor, LocalSerializedTensor]]],
|
||||||
|
load_format: Optional[str] = None,
|
||||||
|
):
|
||||||
|
monkey_patch_torch_reductions()
|
||||||
|
if load_format == "flattened_bucket":
|
||||||
|
# Handle flattened bucket format
|
||||||
|
return self._update_weights_from_flattened_bucket(
|
||||||
|
flattened_tensor_bucket_dict=named_tensors
|
||||||
|
)
|
||||||
|
|
||||||
|
# We need to get device after patch otherwise the device would be wrong
|
||||||
|
device_module = torch.get_device_module(self.device)
|
||||||
|
infered_device = device_module.current_device()
|
||||||
|
|
||||||
|
named_tensors = [
|
||||||
|
(name, _unwrap_tensor(tensor, tp_rank=self.tp_rank, device=infered_device))
|
||||||
|
for name, tensor in named_tensors
|
||||||
|
]
|
||||||
|
if load_format == "direct":
|
||||||
|
_model_load_weights_direct(self.get_model(), named_tensors)
|
||||||
|
elif load_format in self.custom_weight_loaders:
|
||||||
|
custom_loader = dynamic_import(load_format)
|
||||||
|
custom_loader(self.get_model(), named_tensors)
|
||||||
|
elif load_format is None:
|
||||||
|
self.get_model().load_weights(named_tensors)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"Unknown load_format={load_format}")
|
||||||
|
return True, "Success"
|
||||||
|
|
||||||
|
def _update_weights_from_flattened_bucket(
|
||||||
|
self: WeightUpdater,
|
||||||
|
flattened_tensor_bucket_dict,
|
||||||
|
):
|
||||||
|
"""Handle flattened bucket format for weight updates"""
|
||||||
|
flattened_tensor = flattened_tensor_bucket_dict["flattened_tensor"]
|
||||||
|
metadata = flattened_tensor_bucket_dict["metadata"]
|
||||||
|
|
||||||
|
# Convert metadata dict to our format
|
||||||
|
converted_metadata = []
|
||||||
|
for meta in metadata:
|
||||||
|
converted_meta = FlattenedTensorMetadata(
|
||||||
|
name=meta.name,
|
||||||
|
shape=meta.shape,
|
||||||
|
dtype=meta.dtype,
|
||||||
|
start_idx=meta.start_idx,
|
||||||
|
end_idx=meta.end_idx,
|
||||||
|
numel=meta.numel,
|
||||||
|
)
|
||||||
|
converted_metadata.append(converted_meta)
|
||||||
|
|
||||||
|
# Create bucket and reconstruct tensors
|
||||||
|
bucket = FlattenedTensorBucket(
|
||||||
|
flattened_tensor=flattened_tensor, metadata=converted_metadata
|
||||||
|
)
|
||||||
|
reconstructed_tensors = bucket.reconstruct_tensors()
|
||||||
|
|
||||||
|
# Load the reconstructed tensors using the standard method
|
||||||
|
self.get_model().load_weights(reconstructed_tensors)
|
||||||
|
|
||||||
|
return True, "Success"
|
||||||
|
|
||||||
|
def update_weights_from_ipc(self: WeightUpdater, recv_req):
|
||||||
|
"""Update weights from IPC for checkpoint-engine integration."""
|
||||||
|
try:
|
||||||
|
from sglang.srt.checkpoint_engine.checkpoint_engine_worker import (
|
||||||
|
SGLangCheckpointEngineWorkerExtensionImpl,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create a worker extension that integrates with SGLang's model
|
||||||
|
worker = SGLangCheckpointEngineWorkerExtensionImpl(self.get_model_runner())
|
||||||
|
worker.update_weights_from_ipc(recv_req.zmq_handles)
|
||||||
|
return True, "IPC weight update completed successfully"
|
||||||
|
except ImportError as e:
|
||||||
|
return False, f"IPC weight update failed: ImportError {e}"
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"IPC weight update failed: {e}")
|
||||||
|
return False, str(e)
|
||||||
|
|
||||||
|
|
||||||
|
def _model_load_weights_direct(model, named_tensors: List[Tuple[str, torch.Tensor]]):
|
||||||
|
params_dict = dict(model.named_parameters())
|
||||||
|
for name, tensor in named_tensors:
|
||||||
|
default_weight_loader(params_dict[name], tensor)
|
||||||
|
|
||||||
|
|
||||||
|
def _unwrap_tensor(tensor, tp_rank, device):
|
||||||
|
if isinstance(tensor, LocalSerializedTensor):
|
||||||
|
tensor = tensor.get(tp_rank)
|
||||||
|
return tensor.to(device)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LocalSerializedTensor:
|
||||||
|
"""torch.Tensor that gets serialized by MultiprocessingSerializer (which only serializes a pointer and not the data).
|
||||||
|
The i-th element in the list corresponds to i-th rank's GPU."""
|
||||||
|
|
||||||
|
values: List[bytes]
|
||||||
|
|
||||||
|
def get(self, rank: int):
|
||||||
|
return MultiprocessingSerializer.deserialize(self.values[rank])
|
||||||
@@ -1738,18 +1738,22 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||||
success, message = self._draft_worker.draft_runner.update_weights_from_disk(
|
success, message = (
|
||||||
recv_req.model_path,
|
self._draft_worker.draft_runner.weight_updater.update_weights_from_disk(
|
||||||
recv_req.load_format,
|
recv_req.model_path,
|
||||||
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
recv_req.load_format,
|
||||||
|
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
if not success:
|
if not success:
|
||||||
return success, message
|
return success, message
|
||||||
return True, "Succeeded to update model weights."
|
return True, "Succeeded to update model weights."
|
||||||
|
|
||||||
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
||||||
success, message = self._draft_worker.draft_runner.update_weights_from_ipc(
|
success, message = (
|
||||||
recv_req
|
self._draft_worker.draft_runner.weight_updater.update_weights_from_ipc(
|
||||||
|
recv_req
|
||||||
|
)
|
||||||
)
|
)
|
||||||
if not success:
|
if not success:
|
||||||
return success, message
|
return success, message
|
||||||
@@ -1760,15 +1764,19 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
named_tensors = MultiprocessingSerializer.deserialize(
|
named_tensors = MultiprocessingSerializer.deserialize(
|
||||||
recv_req.serialized_named_tensors[self.tp_rank]
|
recv_req.serialized_named_tensors[self.tp_rank]
|
||||||
)
|
)
|
||||||
success, message = self.draft_worker.draft_runner.update_weights_from_tensor(
|
success, message = (
|
||||||
named_tensors=named_tensors,
|
self.draft_worker.draft_runner.weight_updater.update_weights_from_tensor(
|
||||||
load_format=recv_req.load_format,
|
named_tensors=named_tensors,
|
||||||
|
load_format=recv_req.load_format,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
if not success:
|
if not success:
|
||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
success, message = self.target_worker.model_runner.update_weights_from_tensor(
|
success, message = (
|
||||||
named_tensors=named_tensors,
|
self.target_worker.model_runner.weight_updater.update_weights_from_tensor(
|
||||||
load_format=recv_req.load_format,
|
named_tensors=named_tensors,
|
||||||
|
load_format=recv_req.load_format,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
return success, message
|
return success, message
|
||||||
|
|||||||
@@ -884,7 +884,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
for i in range(self.speculative_num_steps):
|
for i in range(self.speculative_num_steps):
|
||||||
success, message = self._draft_worker.draft_runner_list[
|
success, message = self._draft_worker.draft_runner_list[
|
||||||
i
|
i
|
||||||
].update_weights_from_disk(
|
].weight_updater.update_weights_from_disk(
|
||||||
recv_req.model_path,
|
recv_req.model_path,
|
||||||
recv_req.load_format,
|
recv_req.load_format,
|
||||||
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
||||||
@@ -897,7 +897,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
for i in range(self.speculative_num_steps):
|
for i in range(self.speculative_num_steps):
|
||||||
success, message = self._draft_worker.draft_runner_list[
|
success, message = self._draft_worker.draft_runner_list[
|
||||||
i
|
i
|
||||||
].update_weights_from_ipc(recv_req)
|
].weight_updater.update_weights_from_ipc(recv_req)
|
||||||
if not success:
|
if not success:
|
||||||
return success, message
|
return success, message
|
||||||
return True, "Succeeded to update model weights."
|
return True, "Succeeded to update model weights."
|
||||||
|
|||||||
@@ -2748,6 +2748,7 @@ class SafeUnpickler(pickle.Unpickler):
|
|||||||
# --- SGLang & Unitest ---
|
# --- SGLang & Unitest ---
|
||||||
"sglang.srt.weight_sync.tensor_bucket.",
|
"sglang.srt.weight_sync.tensor_bucket.",
|
||||||
"sglang.srt.model_executor.model_runner.",
|
"sglang.srt.model_executor.model_runner.",
|
||||||
|
"sglang.srt.model_executor.model_runner_components.weight_updater.",
|
||||||
"sglang.srt.layers.",
|
"sglang.srt.layers.",
|
||||||
"sglang.srt.utils.",
|
"sglang.srt.utils.",
|
||||||
"sglang.srt.disaggregation.",
|
"sglang.srt.disaggregation.",
|
||||||
|
|||||||
@@ -7,7 +7,9 @@ from torch.distributed.tensor import DTensor
|
|||||||
|
|
||||||
from sglang.srt.entrypoints.engine import Engine
|
from sglang.srt.entrypoints.engine import Engine
|
||||||
from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput
|
from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput
|
||||||
from sglang.srt.model_executor.model_runner import LocalSerializedTensor
|
from sglang.srt.model_executor.model_runner_components.weight_updater import (
|
||||||
|
LocalSerializedTensor,
|
||||||
|
)
|
||||||
from sglang.srt.utils import MultiprocessingSerializer
|
from sglang.srt.utils import MultiprocessingSerializer
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -96,9 +96,7 @@ class TestUpdateWeightsFromTensor(CustomTestCase):
|
|||||||
engine.shutdown()
|
engine.shutdown()
|
||||||
|
|
||||||
def test_update_weights_from_tensor_load_format_custom(self):
|
def test_update_weights_from_tensor_load_format_custom(self):
|
||||||
custom_loader_name = (
|
custom_loader_name = "sglang.srt.model_executor.model_runner_components.weight_updater._model_load_weights_direct"
|
||||||
"sglang.srt.model_executor.model_runner._model_load_weights_direct"
|
|
||||||
)
|
|
||||||
engine = sgl.Engine(
|
engine = sgl.Engine(
|
||||||
model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
custom_weight_loader=[custom_loader_name],
|
custom_weight_loader=[custom_loader_name],
|
||||||
|
|||||||
Reference in New Issue
Block a user