Delegate ModelExpress loading to package (#24723)
Signed-off-by: Zheng Luo <zheluo@nvidia.com>
This commit is contained in:
@@ -79,15 +79,9 @@ class LoadConfig:
|
||||
remote_instance_weight_loader_send_weights_group_ports: Optional[List[int]] = None
|
||||
remote_instance_weight_loader_backend: Optional[str] = None
|
||||
remote_instance_weight_loader_transfer_engine: Optional[Any] = None
|
||||
remote_instance_weight_loader_transfer_engine_session_id: Optional[str] = None
|
||||
modelexpress_url: Optional[str] = None
|
||||
modelexpress_model_name: Optional[str] = None
|
||||
# Fields for building SourceIdentity (needed by both seed and client)
|
||||
modelexpress_tp_size: Optional[int] = None
|
||||
modelexpress_pp_size: Optional[int] = None
|
||||
modelexpress_ep_size: Optional[int] = None
|
||||
modelexpress_dtype: Optional[str] = None
|
||||
modelexpress_quantization: Optional[str] = None
|
||||
modelexpress_transport: str = "transfer_engine"
|
||||
modelexpress_transport: str = "nixl"
|
||||
|
||||
# ModelOpt-specific loading options
|
||||
modelopt_checkpoint_restore_path: Optional[str] = None
|
||||
|
||||
@@ -25,7 +25,6 @@ import os
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
@@ -655,6 +654,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
|
||||
if (
|
||||
self.server_args.remote_instance_weight_loader_use_transfer_engine()
|
||||
# ModelExpress owns TransferEngine memory registration and metadata
|
||||
# publishing for backend=modelexpress. Re-registering here would
|
||||
# overlap the same weight buffers.
|
||||
and self.server_args.remote_instance_weight_loader_backend
|
||||
!= RemoteInstanceWeightLoaderBackend.MODELEXPRESS
|
||||
and self.remote_instance_transfer_engine is not None
|
||||
and self.remote_instance_transfer_engine_weight_info is None
|
||||
):
|
||||
@@ -963,155 +967,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
f"Failed to register transfer engine info for tp_rank={self.tp_rank}: {e}"
|
||||
)
|
||||
|
||||
def _publish_modelexpress_metadata(self):
|
||||
"""Publish metadata to ModelExpress server (seed mode).
|
||||
|
||||
Supports two transport backends:
|
||||
- transfer_engine: publishes TransferEngine session_id (Mooncake)
|
||||
- nixl: creates NIXL agent, registers tensors, publishes nixl_metadata
|
||||
"""
|
||||
try:
|
||||
from modelexpress import p2p_pb2
|
||||
from modelexpress.client import MxClient
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"ModelExpress support requires the 'modelexpress' package. "
|
||||
"Install it with: pip install modelexpress"
|
||||
) from exc
|
||||
|
||||
model_name = (
|
||||
self.server_args.modelexpress_model_name or self.server_args.model_path
|
||||
)
|
||||
mx_url = self.server_args.modelexpress_url
|
||||
transport = self.server_args.modelexpress_transport
|
||||
|
||||
# Build SourceIdentity for this instance
|
||||
identity = p2p_pb2.SourceIdentity(
|
||||
model_name=model_name,
|
||||
backend_framework=p2p_pb2.BACKEND_FRAMEWORK_SGLANG,
|
||||
tensor_parallel_size=self.server_args.tp_size,
|
||||
pipeline_parallel_size=self.server_args.pp_size,
|
||||
expert_parallel_size=self.server_args.ep_size,
|
||||
dtype=self.server_args.dtype or "",
|
||||
quantization=self.server_args.quantization or "",
|
||||
)
|
||||
|
||||
if transport == "nixl":
|
||||
worker, tensor_count = self._build_nixl_worker_metadata(p2p_pb2)
|
||||
else:
|
||||
worker, tensor_count = self._build_transfer_engine_worker_metadata(p2p_pb2)
|
||||
if worker is None:
|
||||
return
|
||||
|
||||
# Generate a unique worker_id for this running instance
|
||||
worker_id = str(uuid.uuid4())
|
||||
|
||||
mx_client = MxClient(server_url=mx_url)
|
||||
try:
|
||||
logger.info(
|
||||
"ModelExpress source [%s]: publishing metadata for model=%s, "
|
||||
"tp_rank=%d, %d tensors, worker_id=%s",
|
||||
transport,
|
||||
model_name,
|
||||
self.tp_rank,
|
||||
tensor_count,
|
||||
worker_id,
|
||||
)
|
||||
mx_source_id = mx_client.publish_metadata(identity, worker, worker_id)
|
||||
mx_client.update_status(
|
||||
mx_source_id=mx_source_id,
|
||||
worker_id=worker_id,
|
||||
worker_rank=self.tp_rank,
|
||||
status=p2p_pb2.SOURCE_STATUS_READY,
|
||||
)
|
||||
logger.info(
|
||||
"ModelExpress source: published ready for model=%s, "
|
||||
"tp_rank=%d, mx_source_id=%s",
|
||||
model_name,
|
||||
self.tp_rank,
|
||||
mx_source_id,
|
||||
)
|
||||
finally:
|
||||
mx_client.close()
|
||||
|
||||
def _build_transfer_engine_worker_metadata(self, p2p_pb2):
|
||||
"""Build WorkerMetadata using TransferEngine session_id."""
|
||||
session_id = self.remote_instance_transfer_engine_session_id
|
||||
weight_info = self.remote_instance_transfer_engine_weight_info
|
||||
|
||||
if not session_id or weight_info is None:
|
||||
logger.warning(
|
||||
"ModelExpress source: skipping publish -- "
|
||||
"TransferEngine not initialized or no weight info"
|
||||
)
|
||||
return None, 0
|
||||
|
||||
tensors = []
|
||||
for name, (addr, numel, element_size) in weight_info.items():
|
||||
tensors.append(
|
||||
p2p_pb2.TensorDescriptor(
|
||||
name=name,
|
||||
addr=addr,
|
||||
size=numel * element_size,
|
||||
device_id=self.gpu_id,
|
||||
)
|
||||
)
|
||||
|
||||
worker = p2p_pb2.WorkerMetadata(
|
||||
worker_rank=self.tp_rank,
|
||||
transfer_engine_session_id=session_id,
|
||||
tensors=tensors,
|
||||
)
|
||||
return worker, len(tensors)
|
||||
|
||||
def _build_nixl_worker_metadata(self, p2p_pb2):
|
||||
"""Build WorkerMetadata using NIXL agent for RDMA transfers."""
|
||||
from modelexpress.nixl_transfer import NixlTransferManager
|
||||
|
||||
agent_name = f"sglang-seed-rank{self.tp_rank}-{uuid.uuid4().hex[:8]}"
|
||||
nixl_mgr = NixlTransferManager(agent_name, self.gpu_id)
|
||||
nixl_mgr.initialize()
|
||||
|
||||
# Collect model tensors for NIXL registration
|
||||
model_tensors = {}
|
||||
for name, param in self.model.named_parameters():
|
||||
t = param.data
|
||||
if t.is_contiguous():
|
||||
model_tensors[name] = t
|
||||
else:
|
||||
# Non-contiguous tensors: register underlying storage as byte view
|
||||
sv = torch.empty(0, dtype=torch.uint8, device=t.device).set_(
|
||||
t.untyped_storage()
|
||||
)
|
||||
if sv.data_ptr() not in {v.data_ptr() for v in model_tensors.values()}:
|
||||
model_tensors[f"{name}.__storage"] = sv
|
||||
|
||||
nixl_metadata = nixl_mgr.register_tensors(model_tensors)
|
||||
|
||||
# Build tensor descriptors from registered tensors
|
||||
tensors = []
|
||||
for td in nixl_mgr.tensor_descriptors:
|
||||
tensors.append(
|
||||
p2p_pb2.TensorDescriptor(
|
||||
name=td.name,
|
||||
addr=td.addr,
|
||||
size=td.size,
|
||||
device_id=td.device_id,
|
||||
dtype=td.dtype,
|
||||
)
|
||||
)
|
||||
|
||||
worker = p2p_pb2.WorkerMetadata(
|
||||
worker_rank=self.tp_rank,
|
||||
nixl_metadata=nixl_metadata,
|
||||
tensors=tensors,
|
||||
)
|
||||
|
||||
# Keep reference alive so NIXL agent isn't garbage collected
|
||||
self._nixl_manager = nixl_mgr
|
||||
|
||||
return worker, len(tensors)
|
||||
|
||||
def model_specific_adjustment(self):
|
||||
server_args = self.server_args
|
||||
|
||||
@@ -1399,14 +1254,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
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,
|
||||
remote_instance_weight_loader_transfer_engine_session_id=self.remote_instance_transfer_engine_session_id,
|
||||
modelexpress_url=self.server_args.modelexpress_url,
|
||||
modelexpress_model_name=self.server_args.modelexpress_model_name
|
||||
or self.server_args.model_path,
|
||||
modelexpress_tp_size=self.server_args.tp_size,
|
||||
modelexpress_pp_size=self.server_args.pp_size,
|
||||
modelexpress_ep_size=self.server_args.ep_size,
|
||||
modelexpress_dtype=self.server_args.dtype,
|
||||
modelexpress_quantization=self.server_args.quantization or "",
|
||||
modelexpress_transport=self.server_args.modelexpress_transport,
|
||||
modelopt_config=modelopt_config,
|
||||
rl_quant_profile=self.server_args.rl_quant_profile,
|
||||
@@ -1464,25 +1313,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
torch.npu.empty_cache()
|
||||
monkey_patch_vllm_parallel_state(reverse=True)
|
||||
|
||||
# Publish metadata to ModelExpress if running as seed source
|
||||
if self.server_args.modelexpress_source:
|
||||
# Seed loads via DefaultModelLoader (load_format=auto), which doesn't
|
||||
# call register_memory_region(). Do it here so weight_info is populated.
|
||||
if (
|
||||
self.remote_instance_transfer_engine_weight_info is None
|
||||
and self.remote_instance_transfer_engine is not None
|
||||
):
|
||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
register_memory_region,
|
||||
)
|
||||
|
||||
self.remote_instance_transfer_engine_weight_info = (
|
||||
register_memory_region(
|
||||
self.model, self.remote_instance_transfer_engine
|
||||
)
|
||||
)
|
||||
self._publish_modelexpress_metadata()
|
||||
|
||||
if not self.is_draft_worker:
|
||||
get_offloader().post_init()
|
||||
|
||||
|
||||
@@ -2200,10 +2200,18 @@ class RemoteInstanceModelLoader(BaseModelLoader):
|
||||
load_config.remote_instance_weight_loader_backend
|
||||
== RemoteInstanceWeightLoaderBackend.MODELEXPRESS
|
||||
):
|
||||
self.load_model_from_modelexpress(
|
||||
model,
|
||||
load_config,
|
||||
device_config,
|
||||
try:
|
||||
from modelexpress.engines.sglang.loader import MxModelLoader
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"ModelExpress support requires the 'modelexpress' "
|
||||
"package. Install it in the SGLang image."
|
||||
) from exc
|
||||
|
||||
model = MxModelLoader(load_config).load_model(
|
||||
model=model,
|
||||
model_config=model_config,
|
||||
device_config=device_config,
|
||||
)
|
||||
else:
|
||||
raise ValueError("Invalid remote instance weight loader backend.")
|
||||
@@ -2316,267 +2324,6 @@ class RemoteInstanceModelLoader(BaseModelLoader):
|
||||
|
||||
return True
|
||||
|
||||
def load_model_from_modelexpress(
|
||||
self,
|
||||
model,
|
||||
load_config: LoadConfig,
|
||||
device_config: DeviceConfig,
|
||||
):
|
||||
"""Load weights via ModelExpress coordination + RDMA transfer.
|
||||
|
||||
Supports two transport backends:
|
||||
- transfer_engine: Mooncake TransferEngine (default)
|
||||
- nixl: NIXL UCX-based RDMA
|
||||
"""
|
||||
try:
|
||||
import grpc
|
||||
from modelexpress import p2p_pb2
|
||||
from modelexpress.client import MxClient
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"ModelExpress support requires the 'modelexpress' package. "
|
||||
"Install it with: pip install modelexpress"
|
||||
) from exc
|
||||
|
||||
tp_rank = load_config.tp_rank
|
||||
model_name = load_config.modelexpress_model_name
|
||||
transport = load_config.modelexpress_transport
|
||||
|
||||
# Process quantized weights to establish final tensor layout
|
||||
target_device = torch.device(device_config.device)
|
||||
for _, module in model.named_modules():
|
||||
quant_method = getattr(module, "quant_method", None)
|
||||
if quant_method is not None:
|
||||
with device_loading_context(module, target_device):
|
||||
quant_method.process_weights_after_loading(module)
|
||||
|
||||
# Register local memory for the chosen transport
|
||||
if transport == "nixl":
|
||||
nixl_mgr = self._init_nixl_for_target(model, load_config, device_config)
|
||||
else:
|
||||
transfer_engine = load_config.remote_instance_weight_loader_transfer_engine
|
||||
if transfer_engine is None:
|
||||
raise RuntimeError(
|
||||
"TransferEngine is not initialized for modelexpress backend."
|
||||
)
|
||||
logger.info(
|
||||
"ModelExpress: registering memory regions for tp_rank=%d...", tp_rank
|
||||
)
|
||||
self.remote_instance_transfer_engine_weight_info = register_memory_region(
|
||||
model, transfer_engine
|
||||
)
|
||||
|
||||
# --- Shared MX discovery logic ---
|
||||
identity = p2p_pb2.SourceIdentity(
|
||||
model_name=model_name,
|
||||
backend_framework=p2p_pb2.BACKEND_FRAMEWORK_SGLANG,
|
||||
tensor_parallel_size=load_config.modelexpress_tp_size or 1,
|
||||
pipeline_parallel_size=load_config.modelexpress_pp_size or 1,
|
||||
expert_parallel_size=load_config.modelexpress_ep_size or 1,
|
||||
dtype=load_config.modelexpress_dtype or "",
|
||||
quantization=load_config.modelexpress_quantization or "",
|
||||
)
|
||||
|
||||
mx_client = MxClient(server_url=load_config.modelexpress_url)
|
||||
try:
|
||||
logger.info(
|
||||
"ModelExpress [%s]: looking for seed (model=%s, rank=%d)...",
|
||||
transport,
|
||||
model_name,
|
||||
tp_rank,
|
||||
)
|
||||
try:
|
||||
resp = mx_client.list_sources(
|
||||
identity=identity,
|
||||
status_filter=p2p_pb2.SOURCE_STATUS_READY,
|
||||
)
|
||||
except grpc.RpcError as e:
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: cannot reach server at "
|
||||
f"{load_config.modelexpress_url}: "
|
||||
f"{e.code()}: {e.details()}"
|
||||
) from e
|
||||
|
||||
source_ref = None
|
||||
for inst in resp.instances:
|
||||
if inst.worker_rank == tp_rank:
|
||||
source_ref = inst
|
||||
break
|
||||
|
||||
if source_ref is None:
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: no READY source found for "
|
||||
f"model={model_name}, rank={tp_rank}. "
|
||||
f"Ensure the seed instance is running and has published metadata."
|
||||
)
|
||||
|
||||
response = mx_client.get_metadata(
|
||||
mx_source_id=source_ref.mx_source_id,
|
||||
worker_id=source_ref.worker_id,
|
||||
)
|
||||
if not response.found:
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: no metadata found for "
|
||||
f"source_id={source_ref.mx_source_id}, "
|
||||
f"worker_id={source_ref.worker_id}"
|
||||
)
|
||||
|
||||
source_worker = response.worker
|
||||
finally:
|
||||
mx_client.close()
|
||||
|
||||
# --- Transport-specific transfer ---
|
||||
if transport == "nixl":
|
||||
self._transfer_via_nixl(model, nixl_mgr, source_worker, tp_rank)
|
||||
else:
|
||||
self._transfer_via_transfer_engine(
|
||||
model, transfer_engine, source_worker, tp_rank
|
||||
)
|
||||
|
||||
_post_load_weights(model)
|
||||
|
||||
logger.info("ModelExpress: weight transfer complete for tp_rank=%d", tp_rank)
|
||||
|
||||
def _transfer_via_transfer_engine(
|
||||
self, model, transfer_engine, source_worker, tp_rank
|
||||
):
|
||||
"""Execute weight transfer using Mooncake TransferEngine."""
|
||||
backend_field = source_worker.WhichOneof("backend_metadata")
|
||||
if backend_field != "transfer_engine_session_id":
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: expected transfer_engine_session_id, "
|
||||
f"got backend_metadata={backend_field}"
|
||||
)
|
||||
seed_session_id = source_worker.transfer_engine_session_id
|
||||
|
||||
seed_weight_info = {}
|
||||
for td in source_worker.tensors:
|
||||
seed_weight_info[td.name] = (td.addr, td.size)
|
||||
|
||||
logger.info(
|
||||
"ModelExpress: got %d tensor descriptors from seed (session=%s)",
|
||||
len(seed_weight_info),
|
||||
seed_session_id,
|
||||
)
|
||||
|
||||
seed_ptr_list = []
|
||||
client_ptr_list = []
|
||||
client_len_list = []
|
||||
for name, tensor in model.named_parameters():
|
||||
weight_info = seed_weight_info.get(name, None)
|
||||
if weight_info is None:
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: cannot find weight info for {name} "
|
||||
f"in seed metadata"
|
||||
)
|
||||
seed_ptr, seed_size = weight_info
|
||||
local_size = tensor.numel() * tensor.element_size()
|
||||
if seed_size != local_size:
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: size mismatch for {name}: "
|
||||
f"seed={seed_size} bytes, local={local_size} bytes"
|
||||
)
|
||||
seed_ptr_list.append(seed_ptr)
|
||||
client_ptr_list.append(tensor.data_ptr())
|
||||
client_len_list.append(local_size)
|
||||
|
||||
logger.info(
|
||||
"ModelExpress: starting TransferEngine RDMA of %d tensors...",
|
||||
len(seed_ptr_list),
|
||||
)
|
||||
ret = transfer_engine.batch_transfer_sync_read(
|
||||
seed_session_id,
|
||||
client_ptr_list,
|
||||
seed_ptr_list,
|
||||
client_len_list,
|
||||
)
|
||||
if ret < 0:
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: batch_transfer_sync_read failed, error={ret}"
|
||||
)
|
||||
|
||||
def _init_nixl_for_target(self, model, load_config, device_config):
|
||||
"""Initialize NIXL agent and register local tensors for the target."""
|
||||
import uuid
|
||||
|
||||
from modelexpress.nixl_transfer import NixlTransferManager
|
||||
|
||||
tp_rank = load_config.tp_rank
|
||||
device_id = device_config.gpu_id
|
||||
|
||||
agent_name = f"sglang-target-rank{tp_rank}-{uuid.uuid4().hex[:8]}"
|
||||
nixl_mgr = NixlTransferManager(agent_name, device_id)
|
||||
nixl_mgr.initialize()
|
||||
|
||||
# Collect local tensors, handling non-contiguous via storage views
|
||||
local_tensors = {}
|
||||
seen_ptrs = set()
|
||||
for name, param in model.named_parameters():
|
||||
t = param.data
|
||||
if t.is_contiguous():
|
||||
ptr = t.data_ptr()
|
||||
if ptr in seen_ptrs:
|
||||
continue
|
||||
seen_ptrs.add(ptr)
|
||||
local_tensors[name] = t
|
||||
else:
|
||||
sv = torch.empty(0, dtype=torch.uint8, device=t.device).set_(
|
||||
t.untyped_storage()
|
||||
)
|
||||
ptr = sv.data_ptr()
|
||||
if ptr in seen_ptrs:
|
||||
continue
|
||||
seen_ptrs.add(ptr)
|
||||
local_tensors[f"{name}.__storage"] = sv
|
||||
|
||||
nixl_mgr.register_tensors(local_tensors)
|
||||
logger.info(
|
||||
"ModelExpress [nixl]: registered %d tensors for tp_rank=%d",
|
||||
len(local_tensors),
|
||||
tp_rank,
|
||||
)
|
||||
return nixl_mgr
|
||||
|
||||
def _transfer_via_nixl(self, model, nixl_mgr, source_worker, tp_rank):
|
||||
"""Execute weight transfer using NIXL RDMA."""
|
||||
from modelexpress.types import TensorDescriptor
|
||||
|
||||
backend_field = source_worker.WhichOneof("backend_metadata")
|
||||
if backend_field != "nixl_metadata":
|
||||
raise RuntimeError(
|
||||
f"ModelExpress: expected nixl_metadata, "
|
||||
f"got backend_metadata={backend_field}"
|
||||
)
|
||||
|
||||
source_tensors = [
|
||||
TensorDescriptor(
|
||||
name=td.name,
|
||||
addr=td.addr,
|
||||
size=td.size,
|
||||
device_id=td.device_id,
|
||||
dtype=td.dtype,
|
||||
)
|
||||
for td in source_worker.tensors
|
||||
]
|
||||
|
||||
logger.info(
|
||||
"ModelExpress [nixl]: starting RDMA transfer of %d tensors...",
|
||||
len(source_tensors),
|
||||
)
|
||||
|
||||
total_bytes, matched, duration = nixl_mgr.receive_from_source(
|
||||
source_metadata=source_worker.nixl_metadata,
|
||||
source_tensors=source_tensors,
|
||||
coalesce_transfers=False,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"ModelExpress [nixl]: transferred %d tensors, " "%.2f GB in %.2fs",
|
||||
matched,
|
||||
total_bytes / 1e9,
|
||||
duration,
|
||||
)
|
||||
|
||||
|
||||
class RemoteModelLoader(BaseModelLoader):
|
||||
"""Model loader that can load Tensors from remote database."""
|
||||
@@ -3281,12 +3028,11 @@ def get_model_loader(
|
||||
return DummyModelLoader(load_config)
|
||||
|
||||
# ModelOptModelLoader's local-copy quantize-and-export workflow doesn't apply
|
||||
# to RUNAI_STREAMER, which streams weights directly from object storage.
|
||||
# RUNAI_STREAMER loads always fall through to the unconditional branch at
|
||||
# the bottom of this function. This also avoids calling _is_already_quantized()
|
||||
# on RunAI streamer cache paths, where huggingface_hub raises HFValidationError.
|
||||
model_optloader_allowed = (
|
||||
model_config and load_config.load_format != LoadFormat.RUNAI_STREAMER
|
||||
# to non-local loaders. These loaders own their weight transport path and still
|
||||
# initialize the model with ModelOpt quantization config where applicable.
|
||||
model_optloader_allowed = model_config and load_config.load_format not in (
|
||||
LoadFormat.RUNAI_STREAMER,
|
||||
LoadFormat.REMOTE_INSTANCE,
|
||||
)
|
||||
|
||||
if model_optloader_allowed and (
|
||||
|
||||
@@ -3539,19 +3539,7 @@ class ServerArgs:
|
||||
self.custom_weight_loader = []
|
||||
|
||||
if self.load_format == "remote_instance":
|
||||
if self.remote_instance_weight_loader_backend == "modelexpress":
|
||||
# ModelExpress backend: requires url in --modelexpress-config
|
||||
if self.modelexpress_url is None:
|
||||
logger.warning(
|
||||
"Fallback load_format to 'auto' due to missing 'url' in --modelexpress-config."
|
||||
)
|
||||
self.load_format = "auto"
|
||||
elif not self.validate_transfer_engine():
|
||||
logger.warning(
|
||||
"Fallback load_format to 'auto' due to 'transfer_engine' (required by modelexpress) not being supported."
|
||||
)
|
||||
self.load_format = "auto"
|
||||
elif (
|
||||
if self.remote_instance_weight_loader_backend != "modelexpress" and (
|
||||
self.remote_instance_weight_loader_seed_instance_ip is None
|
||||
or self.remote_instance_weight_loader_seed_instance_service_port is None
|
||||
):
|
||||
@@ -6623,7 +6611,7 @@ class ServerArgs:
|
||||
"--modelexpress-config",
|
||||
type=str,
|
||||
default=ServerArgs.modelexpress_config,
|
||||
help='JSON config for ModelExpress P2P weight loading. Keys: "url" (required, gRPC host:port), "model_name" (optional, defaults to --model-path), "source" (optional bool, true for seed mode). Example: \'{"url": "localhost:8001", "model_name": "my-model", "source": true}\'',
|
||||
help='JSON config for ModelExpress P2P weight loading. Keys: "url" (optional gRPC host:port override), "transport" ("nixl" or "transfer_engine"). Example: \'{"url": "localhost:8001", "transport": "nixl"}\'',
|
||||
)
|
||||
|
||||
# For PD-Multiplexing
|
||||
@@ -7267,35 +7255,22 @@ class ServerArgs:
|
||||
def modelexpress_url(self) -> Optional[str]:
|
||||
return self._parsed_modelexpress_config.get("url")
|
||||
|
||||
@property
|
||||
def modelexpress_model_name(self) -> Optional[str]:
|
||||
return self._parsed_modelexpress_config.get("model_name")
|
||||
|
||||
@property
|
||||
def modelexpress_source(self) -> bool:
|
||||
return self._parsed_modelexpress_config.get("source", False)
|
||||
|
||||
@property
|
||||
def modelexpress_transport(self) -> str:
|
||||
"""Transport backend for modelexpress: 'transfer_engine' (default) or 'nixl'."""
|
||||
return self._parsed_modelexpress_config.get("transport", "transfer_engine")
|
||||
"""Transport backend for modelexpress."""
|
||||
return self._parsed_modelexpress_config.get("transport", "nixl")
|
||||
|
||||
def remote_instance_weight_loader_use_transfer_engine(self):
|
||||
# Use TransferEngine as seed backend.
|
||||
if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
|
||||
return True
|
||||
# ModelExpress source mode needs TransferEngine init only if transport is transfer_engine.
|
||||
if (
|
||||
self.modelexpress_source
|
||||
and self.modelexpress_transport == "transfer_engine"
|
||||
):
|
||||
return True
|
||||
# Use TransferEngine as client backend.
|
||||
elif (
|
||||
self.load_format == "remote_instance"
|
||||
and self.remote_instance_weight_loader_backend
|
||||
in ("transfer_engine", "modelexpress")
|
||||
and self.modelexpress_transport == "transfer_engine"
|
||||
if self.load_format == "remote_instance" and (
|
||||
self.remote_instance_weight_loader_backend == "transfer_engine"
|
||||
or (
|
||||
self.remote_instance_weight_loader_backend == "modelexpress"
|
||||
and self.modelexpress_transport == "transfer_engine"
|
||||
)
|
||||
):
|
||||
return True
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user