feat: update ModelExpress metadata API to SourceIdentity-based schema (#21222)
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user