feat: update ModelExpress metadata API to SourceIdentity-based schema (#21222)

This commit is contained in:
Zhongdongming Dai
2026-04-10 13:45:05 -07:00
committed by GitHub
parent 6d8330bdb7
commit 4ace144fae
3 changed files with 91 additions and 29 deletions
+6
View File
@@ -79,6 +79,12 @@ class LoadConfig:
remote_instance_weight_loader_transfer_engine: Optional[Any] = None remote_instance_weight_loader_transfer_engine: Optional[Any] = None
modelexpress_url: Optional[str] = None modelexpress_url: Optional[str] = None
modelexpress_model_name: 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
# ModelOpt-specific loading options # ModelOpt-specific loading options
modelopt_checkpoint_restore_path: Optional[str] = None modelopt_checkpoint_restore_path: Optional[str] = None
@@ -24,6 +24,7 @@ import os
import socket import socket
import threading import threading
import time import time
import uuid
from collections import defaultdict from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from typing import Callable, List, Optional, Tuple, Union from typing import Callable, List, Optional, Tuple, Union
@@ -839,6 +840,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
) )
return return
# 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 "",
)
# Build tensor descriptors from weight_info dict # Build tensor descriptors from weight_info dict
tensors = [] tensors = []
for name, (addr, numel, element_size) in weight_info.items(): for name, (addr, numel, element_size) in weight_info.items():
@@ -857,27 +869,33 @@ class ModelRunner(ModelRunnerKVCacheMixin):
tensors=tensors, tensors=tensors,
) )
# Generate a unique worker_id for this running instance
worker_id = str(uuid.uuid4())
mx_client = MxClient(server_url=mx_url) mx_client = MxClient(server_url=mx_url)
try: try:
logger.info( logger.info(
"ModelExpress source: publishing metadata for model=%s, " "ModelExpress source: publishing metadata for model=%s, "
"tp_rank=%d, session=%s, %d tensors", "tp_rank=%d, session=%s, %d tensors, worker_id=%s",
model_name, model_name,
self.tp_rank, self.tp_rank,
session_id, session_id,
len(tensors), len(tensors),
worker_id,
) )
mx_client.publish_metadata(model_name, [worker]) mx_source_id = mx_client.publish_metadata(identity, worker, worker_id)
mx_client.publish_ready( mx_client.update_status(
model_name, mx_source_id=mx_source_id,
worker_id=self.tp_rank, worker_id=worker_id,
session_id=mx_client.session_id, worker_rank=self.tp_rank,
metadata_hash="", status=p2p_pb2.SOURCE_STATUS_READY,
) )
logger.info( logger.info(
"ModelExpress source: published ready for model=%s, tp_rank=%d", "ModelExpress source: published ready for model=%s, "
"tp_rank=%d, mx_source_id=%s",
model_name, model_name,
self.tp_rank, self.tp_rank,
mx_source_id,
) )
finally: finally:
mx_client.close() mx_client.close()
@@ -1173,6 +1191,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
modelexpress_url=self.server_args.modelexpress_url, modelexpress_url=self.server_args.modelexpress_url,
modelexpress_model_name=self.server_args.modelexpress_model_name modelexpress_model_name=self.server_args.modelexpress_model_name
or self.server_args.model_path, 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 "",
modelopt_config=modelopt_config, modelopt_config=modelopt_config,
rl_quant_profile=self.server_args.rl_quant_profile, rl_quant_profile=self.server_args.rl_quant_profile,
draft_model_idx=self.draft_model_idx, draft_model_idx=self.draft_model_idx,
+54 -21
View File
@@ -2289,6 +2289,8 @@ class RemoteInstanceModelLoader(BaseModelLoader):
): ):
"""Load weights via ModelExpress coordination + TransferEngine RDMA.""" """Load weights via ModelExpress coordination + TransferEngine RDMA."""
try: try:
import grpc
from modelexpress import p2p_pb2
from modelexpress.client import MxClient from modelexpress.client import MxClient
except ImportError as exc: except ImportError as exc:
raise ImportError( raise ImportError(
@@ -2304,6 +2306,13 @@ class RemoteInstanceModelLoader(BaseModelLoader):
tp_rank = load_config.tp_rank tp_rank = load_config.tp_rank
model_name = load_config.modelexpress_model_name model_name = load_config.modelexpress_model_name
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)
logger.info( logger.info(
"ModelExpress: registering memory regions for tp_rank=%d...", tp_rank "ModelExpress: registering memory regions for tp_rank=%d...", tp_rank
) )
@@ -2311,39 +2320,63 @@ class RemoteInstanceModelLoader(BaseModelLoader):
model, transfer_engine model, transfer_engine
) )
# Wait for seed to be ready via ModelExpress # Build SourceIdentity matching the seed's identity
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 "",
)
# Query MX server for a READY source matching our identity and rank
mx_client = MxClient(server_url=load_config.modelexpress_url) mx_client = MxClient(server_url=load_config.modelexpress_url)
try: try:
logger.info( logger.info(
"ModelExpress: waiting for seed ready (model=%s)...", "ModelExpress: looking for seed (model=%s, rank=%d)...",
model_name, model_name,
tp_rank,
) )
ready, session_id, metadata_hash = mx_client.wait_for_ready( try:
model_name, resp = mx_client.list_sources(
worker_id=tp_rank, identity=identity,
) status_filter=p2p_pb2.SOURCE_STATUS_READY,
if not ready: )
except grpc.RpcError as e:
raise RuntimeError( raise RuntimeError(
f"ModelExpress: timed out waiting for seed ready " f"ModelExpress: cannot reach server at "
f"(model={model_name}, worker={tp_rank})" 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(model_name) # Fetch full metadata for the discovered worker
response = mx_client.get_metadata(
mx_source_id=source_ref.mx_source_id,
worker_id=source_ref.worker_id,
)
if not response.found: if not response.found:
raise RuntimeError( raise RuntimeError(
f"ModelExpress: no metadata found for model={model_name}" f"ModelExpress: no metadata found for "
f"source_id={source_ref.mx_source_id}, "
f"worker_id={source_ref.worker_id}"
) )
# Find the worker matching our tp_rank source_worker = response.worker
source_worker = None
for w in response.workers:
if w.worker_rank == tp_rank:
source_worker = w
break
if source_worker is None:
raise RuntimeError(
f"ModelExpress: no worker metadata for rank={tp_rank}"
)
# Extract session_id from oneof backend_metadata # Extract session_id from oneof backend_metadata
backend_field = source_worker.WhichOneof("backend_metadata") backend_field = source_worker.WhichOneof("backend_metadata")