support non disturbing remote instance weight loader v2 (#14997)
Signed-off-by: Anqi Shen <amy.saq@antgroup.com>
This commit is contained in:
@@ -136,9 +136,10 @@ from sglang.srt.model_executor.input_buffers import GraphInputBuffers
|
||||
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||
PiecewiseCudaGraphRunner,
|
||||
)
|
||||
from sglang.srt.model_loader import get_model
|
||||
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
|
||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
RemoteInstanceWeightLoaderBackend,
|
||||
register_memory_region,
|
||||
trigger_init_weights_send_group_for_remote_instance_request,
|
||||
)
|
||||
from sglang.srt.model_loader.utils import set_default_torch_dtype
|
||||
@@ -158,6 +159,7 @@ from sglang.srt.utils import (
|
||||
get_available_gpu_memory,
|
||||
get_bool_env_var,
|
||||
get_cpu_ids_by_node,
|
||||
get_local_ip_auto,
|
||||
init_custom_process_group,
|
||||
is_cuda,
|
||||
is_float4_e2m1fn_x2,
|
||||
@@ -319,6 +321,10 @@ class ModelRunner:
|
||||
self.forward_pass_id = 0
|
||||
self.init_new_workspace = False
|
||||
|
||||
self.remote_instance_transfer_engine = None
|
||||
self.remote_instance_transfer_engine_session_id = ""
|
||||
self.remote_instance_transfer_engine_weight_info = None
|
||||
|
||||
# Apply the rank zero filter to logger
|
||||
if server_args.show_time_cost:
|
||||
enable_show_time_cost()
|
||||
@@ -393,6 +399,9 @@ class ModelRunner:
|
||||
enable=self.server_args.enable_memory_saver
|
||||
)
|
||||
|
||||
if self.server_args.remote_instance_weight_loader_use_transfer_engine():
|
||||
self.remote_instance_init_transfer_engine()
|
||||
|
||||
if not self.is_draft_worker:
|
||||
set_global_expert_location_metadata(
|
||||
compute_initial_expert_location_metadata(
|
||||
@@ -433,6 +442,15 @@ class ModelRunner:
|
||||
self.sampler = Sampler()
|
||||
self.load_model()
|
||||
|
||||
if (
|
||||
self.server_args.remote_instance_weight_loader_use_transfer_engine()
|
||||
and self.remote_instance_transfer_engine is not None
|
||||
and self.remote_instance_transfer_engine_weight_info is None
|
||||
):
|
||||
self.remote_instance_transfer_engine_weight_info = register_memory_region(
|
||||
self.model, self.remote_instance_transfer_engine
|
||||
)
|
||||
|
||||
# Check if the model is using hybrid SWA
|
||||
if (
|
||||
not self.server_args.disable_hybrid_swa_memory
|
||||
@@ -547,6 +565,23 @@ class ModelRunner:
|
||||
# Initialize piecewise CUDA graph
|
||||
self.init_piecewise_cuda_graphs()
|
||||
|
||||
def remote_instance_init_transfer_engine(self):
|
||||
try:
|
||||
from mooncake.engine import TransferEngine
|
||||
except ImportError as e:
|
||||
logger.warning(
|
||||
"Please install mooncake for using remote instance transfer engine: pip install mooncake"
|
||||
)
|
||||
return
|
||||
self.remote_instance_transfer_engine = TransferEngine()
|
||||
local_ip = get_local_ip_auto()
|
||||
self.remote_instance_transfer_engine.initialize(
|
||||
local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.value
|
||||
)
|
||||
self.remote_instance_transfer_engine_session_id = (
|
||||
f"{local_ip}:{self.remote_instance_transfer_engine.get_rpc_port()}"
|
||||
)
|
||||
|
||||
def model_specific_adjustment(self):
|
||||
server_args = self.server_args
|
||||
|
||||
@@ -764,6 +799,8 @@ class ModelRunner:
|
||||
remote_instance_weight_loader_seed_instance_ip=self.server_args.remote_instance_weight_loader_seed_instance_ip,
|
||||
remote_instance_weight_loader_seed_instance_service_port=self.server_args.remote_instance_weight_loader_seed_instance_service_port,
|
||||
remote_instance_weight_loader_send_weights_group_ports=self.server_args.remote_instance_weight_loader_send_weights_group_ports,
|
||||
remote_instance_weight_loader_backend=self.server_args.remote_instance_weight_loader_backend,
|
||||
remote_instance_weight_loader_transfer_engine=self.remote_instance_transfer_engine,
|
||||
modelopt_config=modelopt_config,
|
||||
rl_quant_profile=self.server_args.rl_quant_profile,
|
||||
)
|
||||
@@ -772,7 +809,11 @@ class ModelRunner:
|
||||
self.model_config, self.load_config, self.tp_size
|
||||
)
|
||||
|
||||
if self.server_args.load_format == LoadFormat.REMOTE_INSTANCE:
|
||||
if (
|
||||
self.server_args.load_format == LoadFormat.REMOTE_INSTANCE
|
||||
and self.server_args.remote_instance_weight_loader_backend
|
||||
== RemoteInstanceWeightLoaderBackend.NCCL
|
||||
):
|
||||
if self.tp_rank == 0:
|
||||
instance_ip = socket.gethostbyname(socket.gethostname())
|
||||
t = threading.Thread(
|
||||
@@ -797,11 +838,18 @@ class ModelRunner:
|
||||
GPU_MEMORY_TYPE_WEIGHTS,
|
||||
enable_cpu_backup=enable_cpu_backup,
|
||||
):
|
||||
self.model = get_model(
|
||||
model_config=self.model_config,
|
||||
self.loader = get_model_loader(
|
||||
load_config=self.load_config,
|
||||
model_config=self.model_config,
|
||||
)
|
||||
self.model = self.loader.load_model(
|
||||
model_config=self.model_config,
|
||||
device_config=DeviceConfig(self.device, self.gpu_id),
|
||||
)
|
||||
if hasattr(self.loader, "remote_instance_transfer_engine_weight_info"):
|
||||
self.remote_instance_transfer_engine_weight_info = (
|
||||
self.loader.remote_instance_transfer_engine_weight_info
|
||||
)
|
||||
monkey_patch_vllm_parallel_state(reverse=True)
|
||||
|
||||
get_offloader().post_init()
|
||||
|
||||
Reference in New Issue
Block a user