diff --git a/tensorrt_llm/_torch/models/checkpoints/auto_mapper.py b/tensorrt_llm/_torch/models/checkpoints/auto_mapper.py index 5fb87c349d32..8644bd90de9c 100644 --- a/tensorrt_llm/_torch/models/checkpoints/auto_mapper.py +++ b/tensorrt_llm/_torch/models/checkpoints/auto_mapper.py @@ -11,10 +11,10 @@ def get(format: str, name: Optional[str] = None) -> "BaseWeightMapper": try: return MODEL_CLASS_MAPPER_MAPPING[f'{name}_{format}']() except KeyError: # no mapper for this model architecture, resort to default - if format == "MX": - # MX uses HF on-disk checkpoint format for fallback, so - # an architecture-specific HF mapper is closer than the - # generic MX/HF default mapper. + if format == "modelexpress": + # ModelExpress uses HF on-disk checkpoint format for + # fallback, so an architecture-specific HF mapper is closer + # than the generic ModelExpress/HF default mapper. try: return MODEL_CLASS_MAPPER_MAPPING[f'{name}_HF']() except KeyError: diff --git a/tensorrt_llm/_torch/models/checkpoints/hf/config_loader.py b/tensorrt_llm/_torch/models/checkpoints/hf/config_loader.py index 8bd582211f79..44a554a9531a 100644 --- a/tensorrt_llm/_torch/models/checkpoints/hf/config_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/hf/config_loader.py @@ -7,7 +7,7 @@ from tensorrt_llm._torch.models.modeling_utils import register_config_loader -@register_config_loader("MX") +@register_config_loader("modelexpress") @register_config_loader("HF") class HfConfigLoader(BaseConfigLoader): diff --git a/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py b/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py index 264104c6f2e1..2b47e06f60e7 100644 --- a/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py @@ -34,7 +34,7 @@ from tensorrt_llm.mapping import Mapping -@register_checkpoint_weight_loader("MX") +@register_checkpoint_weight_loader("modelexpress") @register_checkpoint_weight_loader("mistral") @register_checkpoint_weight_loader("mistral_large_3") @register_checkpoint_weight_loader("HF") diff --git a/tensorrt_llm/_torch/models/checkpoints/hf/weight_mapper.py b/tensorrt_llm/_torch/models/checkpoints/hf/weight_mapper.py index 0828d10f6e66..bd328bbfe13b 100644 --- a/tensorrt_llm/_torch/models/checkpoints/hf/weight_mapper.py +++ b/tensorrt_llm/_torch/models/checkpoints/hf/weight_mapper.py @@ -6,7 +6,7 @@ from ..base_weight_mapper import BaseWeightMapper -@register_mapper("MX") +@register_mapper("modelexpress") @register_mapper("HF") class HfWeightMapper(BaseWeightMapper): diff --git a/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py index 304059794a3b..fdbe7bec08a2 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py @@ -13,455 +13,29 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""MX (ModelExpress) checkpoint loader. +"""Compatibility shim for the ModelExpress TRT-LLM checkpoint loader.""" -Thin adapter on top of the upstream ``modelexpress`` Python client -(``ai-dynamo/modelexpress``). All NIXL/RDMA mechanics (agent setup, -tensor registration, source-target name matching, dtype-cast handling, -PVC fallback, etc.) live in the upstream ``MxLiveWeightLoader`` and -``publish_model_params`` helpers — we only call them at the right -points in TRT-LLM's loading lifecycle. +from __future__ import annotations -When no MX server is reachable (or the upstream library is not -installed), this loader transparently falls back to standard -HuggingFace checkpoint loading (disk -> CPU -> GPU) by way of its -``HfCheckpointLoader`` base class. -""" +from typing import Any -import os -import threading -import traceback -from contextlib import contextmanager -from pathlib import Path -from typing import Any, Callable, Optional, Type, Union - -import grpc - -from tensorrt_llm._torch.models.checkpoints.base_config_loader import BaseConfigLoader -from tensorrt_llm._torch.models.checkpoints.base_weight_loader import BaseWeightLoader -from tensorrt_llm._torch.models.checkpoints.base_weight_mapper import BaseWeightMapper -from tensorrt_llm._torch.models.checkpoints.hf.checkpoint_loader import HfCheckpointLoader from tensorrt_llm._torch.models.modeling_utils import register_checkpoint_loader -from tensorrt_llm.logger import logger -from tensorrt_llm.mapping import Mapping - -# Defensive default for the upstream ``MX_SOURCE_QUERY_TIMEOUT`` env var. -# The upstream ``MxLiveWeightLoader`` polls the MX server every 5 s for up -# to ``MX_SOURCE_QUERY_TIMEOUT`` seconds (default 3600 = 1 hour) waiting -# for a source. On a cold cluster (no donor up yet), this means the very -# first replica blocks for an hour before falling back to disk. We cap -# the default at 30 s so first-replica startup degrades gracefully; users -# can still override via the env var or a future per-loader knob. -# Tracked as MX-4 in §15 (non-blocking source-query API upstream). -_MX_SOURCE_QUERY_TIMEOUT_DEFAULT_S = "30" -_MX_PUBLISH_ENV_LOCK = threading.Lock() - - -@contextmanager -def _temporary_env(key: str, value: Optional[str]): - """Temporarily set or clear one environment variable.""" - if value is None: - yield - return - prior = os.environ.get(key) - os.environ[key] = value - try: - yield - finally: - if prior is None: - os.environ.pop(key, None) - else: - os.environ[key] = prior - - -@register_checkpoint_loader("MX") -class MXCheckpointLoader(HfCheckpointLoader): - """Checkpoint loader for MX (ModelExpress) P2P weight transfer. - - When an MX server is reachable AND the upstream ``modelexpress`` - library is installed, weights are transferred directly from a - source instance via NIXL/RDMA, bypassing disk I/O. The source - publishes its weights *before* ``post_load_weights()`` runs so - targets receive raw loaded state and can run their own - post-load transforms. - - When the MX server or library is unavailable, this loader - transparently falls back to standard HuggingFace checkpoint - loading via the parent ``HfCheckpointLoader``. - - All transport-level mechanics (NIXL, dtype casts, source matching, - fallback) are delegated to ``modelexpress.trtllm_live_transfer`` - so that this class stays a thin adapter — when the MX wire - protocol or transport evolves, only the upstream library needs - to track it. - """ - - def __init__( - self, - *, - weight_loader: Optional[BaseWeightLoader] = None, - weight_mapper: Optional[BaseWeightMapper] = None, - config_loader: Optional[BaseConfigLoader] = None, - mx_server_url: Optional[str] = None, - model_name: Optional[Union[str, Path]] = None, - query_timeout_s: Optional[int] = None, - ): - super().__init__( - weight_loader=weight_loader, - weight_mapper=weight_mapper, - config_loader=config_loader, - ) - # HfCheckpointLoader initializes the backing attribute to "HF". Keep it - # aligned with the property override so legacy/internal code that reads - # _checkpoint_format directly does not see a stale value. - self._checkpoint_format = "MX" - self._mx_server_url = mx_server_url - # ``model_name`` is the human-readable identity to publish/look up - # under on the MX server. Typically the user-supplied - # ``llm_args.model`` (a Hub ID like ``"Qwen/Qwen2.5-72B-Instruct"`` - # or a local path). ``publish_as_source()`` resolves it via - # :func:`_resolve_mx_model_name` (with HF-snapshot path fallback). - self._model_name = str(model_name) if model_name is not None else None - self._query_timeout_s = query_timeout_s - self._p2p_succeeded = False - - @property - def checkpoint_format(self) -> str: - """Override parent's checkpoint_format to return 'MX'.""" - return "MX" - - @property - def mx_server_url(self) -> Optional[str]: - return self._mx_server_url - - @property - def model_name(self) -> Optional[str]: - """Explicit model identity passed to the constructor (if any). - - Note this is the *as-configured* value (e.g. ``llm_args.model``), - not the final resolved identity that ends up in the published - ``MODEL_NAME``. The full resolution (with env var and basename - fallbacks) happens inside :meth:`publish_as_source`. - """ - return self._model_name - - @property - def query_timeout_s(self) -> Optional[int]: - return self._query_timeout_s - - @property - def p2p_succeeded(self) -> bool: - """Whether the last load_weights() call used P2P transfer. - - ``True`` means weights are already in model parameter buffers - and the standard weight-mapping pipeline should be skipped - for those parameters. - """ - return self._p2p_succeeded - - def is_weights_preloaded(self) -> bool: - """Whether the last MX load wrote weights directly into the model.""" - return self._p2p_succeeded - - def load_weights(self, checkpoint_dir: str, mapping: Mapping, **kwargs) -> dict[str, Any]: - """Load weights, preferring MX P2P transfer when available. - - Delegates the actual transfer to the upstream - ``modelexpress.trtllm_live_transfer.MxLiveWeightLoader``, - which handles NIXL setup, source discovery, name matching, - dtype casting, and PVC fallback for size-mismatched tensors. - - Args: - checkpoint_dir: Path to the HF checkpoint directory. - mapping: Distributed mapping configuration. - **kwargs: Additional keyword arguments. When ``model`` is - passed it is used as the target for direct P2P writes. - - Returns: - A weights dict. Empty when MX P2P fully succeeded (weights - already in model params); populated when falling back to - disk loading for some or all weights. - """ - model = kwargs.pop("model", None) - self._p2p_succeeded = False - - if self._mx_server_url is None or model is None: - return self._fallback_to_disk( - checkpoint_dir, - mapping, - reason=( - "no MX server URL configured" - if self._mx_server_url is None - else "no model reference passed (cannot do P2P writes)" - ), - **kwargs, - ) - - try: - from modelexpress.trtllm_live_transfer import ( # type: ignore[import-not-found] - MxClient, - MxLiveWeightLoader, - _build_trtllm_identity, - ) - except ImportError: - logger.warning( - "modelexpress library not installed; cannot use MX P2P " - "weight transfer. Install from " - "https://github.com/ai-dynamo/modelexpress (Python client at " - "modelexpress_client/python). Falling back to disk loading." - ) - return self._fallback_to_disk(checkpoint_dir, mapping, **kwargs) - timeout_override = self._resolve_query_timeout_override( - checkpoint_dir, - MxClient, - _build_trtllm_identity, - ) - with _temporary_env("MX_SOURCE_QUERY_TIMEOUT", timeout_override): - try: - mx_loader = MxLiveWeightLoader(mx_server=self._mx_server_url) - fallback_weights = mx_loader.load_weights( - checkpoint_dir, - mapping=mapping, - model=model, - ) - except Exception: - # Deliberately broad: MX is an opportunistic fast path and HF - # disk loading remains the correctness path. Preserve the full - # traceback so unexpected upstream failures are diagnosable. - logger.warning( - f"MX P2P transfer failed; falling back to disk loading.\n{traceback.format_exc()}" - ) - return self._fallback_to_disk(checkpoint_dir, mapping, **kwargs) - if fallback_weights: - fallback_bytes = sum( - tensor.numel() * tensor.element_size() for tensor in fallback_weights.values() - ) - # Mixed-success case: MX delivered matched tensors into model - # params via P2P and returned only size-mismatched tensors for - # the standard disk path to apply. Keep the P2P transfer and - # let ModelLoader merge these fallback tensors. - logger.warning( - "MX P2P returned %d fallback weights (%.2f MiB, size mismatch) " - "from %s. Merging fallback weights through the disk pipeline; " - "if this warning persists for this model, disable MX for it to " - "avoid paying both P2P and disk-loading costs.", - len(fallback_weights), - fallback_bytes / (1 << 20), - self._mx_server_url, - ) - self._p2p_succeeded = True - return fallback_weights - - self._p2p_succeeded = True - logger.info( - "MX P2P weight transfer succeeded from %s", - self._mx_server_url, - ) - return {} - - def _resolve_query_timeout_override( - self, checkpoint_dir: str, MxClient: Type[Any], build_identity: Callable[..., Any] - ) -> Optional[str]: - """Return temporary ``MX_SOURCE_QUERY_TIMEOUT`` override, if any.""" - if self._query_timeout_s is not None: - return str(self._query_timeout_s) - - if os.environ.get("MX_SOURCE_QUERY_TIMEOUT"): - return None - - if self._has_any_source_instance(checkpoint_dir, MxClient, build_identity): - return None - - logger.warning( - f"No MX source is currently registered for {self._resolve_publish_name(checkpoint_dir)}; " - f"using MX_SOURCE_QUERY_TIMEOUT={_MX_SOURCE_QUERY_TIMEOUT_DEFAULT_S} " - "for fast disk fallback. Set mx_config.server_query_timeout_s or " - "MX_SOURCE_QUERY_TIMEOUT for long-running donor-load deployments." - ) - return _MX_SOURCE_QUERY_TIMEOUT_DEFAULT_S - - def _has_any_source_instance( - self, checkpoint_dir: str, MxClient: Type[Any], build_identity: Callable[..., Any] - ) -> bool: - """Best-effort fast probe for registered MX source instances.""" - client = None - try: - identity = build_identity(model_name=self._resolve_publish_name(checkpoint_dir)) - client = MxClient(server_url=self._mx_server_url) - list_resp = client.list_sources(identity=identity) - return bool(getattr(list_resp, "instances", [])) - except (AttributeError, RuntimeError, TimeoutError, grpc.RpcError): - # If the probe cannot complete, prefer fast fallback over the - # upstream 1-hour default. The actual MxLiveWeightLoader call below - # remains the source of truth and may still succeed. - logger.warning( - f"MX source probe failed; using fast fallback timeout.\n{traceback.format_exc()}" - ) - return False - finally: - if client is not None and hasattr(client, "close"): - client.close() - - def _resolve_publish_name(self, checkpoint_dir: Optional[str]) -> str: - return _resolve_mx_model_name(self._model_name, checkpoint_dir) - - def _fallback_to_disk( - self, checkpoint_dir: str, mapping: Mapping, *, reason: Optional[str] = None, **kwargs - ) -> dict[str, Any]: - """Standard HF disk loading fallback.""" - if reason is not None: - logger.info(f"MX P2P unavailable ({reason}); loading from disk: {checkpoint_dir}") - else: - logger.info(f"MX P2P unavailable; loading from disk: {checkpoint_dir}") - return super().load_weights(checkpoint_dir, mapping=mapping, **kwargs) - - def publish_as_source( - self, - model, - checkpoint_dir: Optional[str] = None, - ) -> None: - """Publish this instance's weights so other ranks can pull via P2P. - - Called by the integration in ``model_loader.py`` *before* - ``post_load_weights()`` so targets receive raw loaded state and - can apply their own post-load transforms. - - Delegates to the upstream - ``modelexpress.trtllm_live_transfer.publish_model_params`` - helper, which handles the per-rank NIXL setup, tensor - registration, and gRPC publish. - - Args: - model: The model whose weights to publish. - checkpoint_dir: Checkpoint directory. Used as a last-resort - fallback for resolving the ``MODEL_NAME`` identity when - neither ``model_name`` was passed to the constructor nor - ``MODEL_NAME`` is set in the environment. - """ - - if self._mx_server_url is None: - return +@register_checkpoint_loader("modelexpress") +class MXCheckpointLoader: + """Instantiate the ModelExpress-owned TRT-LLM checkpoint loader.""" + def __new__(cls, *args: Any, **kwargs: Any): try: - from modelexpress.trtllm_live_transfer import ( # type: ignore[import-not-found] - publish_model_params, - ) - except ImportError: - logger.debug("modelexpress library not installed; skipping MX publish.") - return - - # THREADSAFETY: upstream publish_model_params reads MODEL_EXPRESS_URL and - # MODEL_NAME from the environment. Set both from our resolved - # configuration so per-instance values (URL passed via - # llm_args.mx_config.server_url, identity from llm_args.model) are - # respected, then restore prior state. This is safe for the current - # sequential TRT-LLM worker path, but co-resident ranks in one Python - # interpreter would race on process-wide env. Tracked as MX-2 in §15 - # (the env-var dance goes away when upstream exports a public identity - # builder / publish API). - resolved_name = self._resolve_publish_name(checkpoint_dir) - - env_overrides = { - "MODEL_EXPRESS_URL": self._mx_server_url, - "MODEL_NAME": resolved_name, - } - if threading.active_count() > 1: - logger.warning_once( - "MX publish uses process-wide MODEL_EXPRESS_URL/MODEL_NAME " - "environment variables; concurrent publish calls in one Python " - "process are serialized, but unrelated env readers can still " - "observe transient values. Tracked by MX-2.", - key="mx_publish_env_threaded_warning", - ) - with _MX_PUBLISH_ENV_LOCK: - prior = {key: os.environ.get(key) for key in env_overrides} - for key, value in env_overrides.items(): - os.environ[key] = value - - try: - publish_model_params(model) - logger.info( - "Published weights to MX server at %s as model=%r", - self._mx_server_url, - resolved_name, - ) - except Exception: - # Deliberately broad: publish is best-effort. A publish failure - # should not fail the local worker that already loaded weights. - logger.warning( - f"Failed to publish weights to MX server at {self._mx_server_url}.\n" - f"{traceback.format_exc()}" - ) - finally: - for key, prior_value in prior.items(): - if prior_value is None: - os.environ.pop(key, None) - else: - os.environ[key] = prior_value - - def post_load_publish( - self, model, *, checkpoint_dir: str, weights_preloaded: bool = False - ) -> None: - """Publish only workers that loaded locally, not MX P2P receivers.""" - if weights_preloaded: - return - self.publish_as_source(model, checkpoint_dir=checkpoint_dir) - - -# --------------------------------------------------------------------------- -# Module-level helpers -# --------------------------------------------------------------------------- - - -def _resolve_mx_model_name(model_name_arg: Optional[str], checkpoint_dir: Optional[str]) -> str: - """Resolve a stable model identity for publishing to the MX server. - - Resolution order (first non-empty wins): - - 1. ``model_name_arg`` — the explicit value passed at construction - time (typically ``llm_args.model``: a Hub ID like - ``"Qwen/Qwen2.5-72B-Instruct"`` or a local path). - 2. ``MODEL_NAME`` env var — upstream's existing convention. - 3. ``checkpoint_dir`` basename, with HF-snapshot path fallback so - ``.../models----/snapshots//`` resolves to - ``"/"`` instead of the commit hash. - 4. Literal ``"unknown"`` — matches upstream's own sentinel. - """ - candidate = model_name_arg or os.environ.get("MODEL_NAME") or checkpoint_dir - if not candidate: - return "unknown" - return _normalize_model_identity(str(candidate)) - - -def _normalize_model_identity(s: str) -> str: - """Convert a model identifier to a stable, human-readable name. + from modelexpress.engines.trtllm.loader import MXCheckpointLoader as _MXCheckpointLoader + except ImportError as exc: + raise ImportError( + "checkpoint_format='modelexpress' requires the modelexpress Python " + "package with TRT-LLM support installed." + ) from exc - Hub IDs (``"org/name"``) and arbitrary user-provided strings are - returned unchanged. Filesystem paths are reduced to a basename, with - HuggingFace cache snapshot layouts (``snapshots//``) - walked up to recover the original ``"org/name"`` identity. - """ - if not s: - return "unknown" + return _MXCheckpointLoader(*args, **kwargs) - # Heuristic: a Hub ID is bare ``"name"`` or ``"org/name"``. Anything - # that starts with a path separator/expansion or contains more than - # one "/" is treated as a path. Single-"/" strings remain ambiguous; - # avoid an NFS ``exists`` probe for common Hub IDs and only touch the - # filesystem when the string has explicit local-path syntax. - looks_like_path = s.startswith(("/", "./", "../", "~")) or s.count("/") > 1 - if not looks_like_path: - return s - p = Path(s).expanduser() - name = p.name - if name and "snapshots" in p.parts: - # HF cache layout: ``.../models----/snapshots//``. - # Walk up to find the ``models----`` directory and - # un-mangle it back to ``"/"``. - for ancestor in p.parents: - if ancestor.name.startswith("models--"): - return ancestor.name[len("models--") :].replace("--", "/") - return name or "unknown" +__all__ = ["MXCheckpointLoader"] diff --git a/tensorrt_llm/_torch/pyexecutor/model_loader.py b/tensorrt_llm/_torch/pyexecutor/model_loader.py index 54c02754f12d..132b991983da 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_loader.py +++ b/tensorrt_llm/_torch/pyexecutor/model_loader.py @@ -204,9 +204,9 @@ def _construct_checkpoint_loader( checkpoint_format)() config_loader = get_config_loader(checkpoint_format)() - # Pass extra kwargs for format-specific loaders (e.g. MX). + # Pass extra kwargs for format-specific loaders (e.g. ModelExpress). extra_kwargs: dict = {} - if checkpoint_format == "MX": + if checkpoint_format == "modelexpress": if mx_config is not None: extra_kwargs["mx_server_url"] = mx_config.server_url extra_kwargs[ @@ -418,10 +418,10 @@ def init_meta_tensor(t: torch.Tensor): ) weights_preloaded = False if load_format == LoadFormat.AUTO: - # Pass model= so format-specific loaders (e.g. MX) can + # Pass model= so format-specific loaders (e.g. ModelExpress) can # write weights directly into parameter buffers via P2P. # Generic loaders ignore model=; loaders that can consume a - # live module reference (MX) use it for direct writes. + # live module reference use it for direct writes. load_weights_kwargs: dict = { "mapping": self.mapping, "model": model, @@ -434,7 +434,7 @@ def init_meta_tensor(t: torch.Tensor): weights = checkpoint_loader.load_weights( checkpoint_dir, **load_weights_kwargs) - # When MX P2P succeeds, weights are already in model params. + # When ModelExpress P2P succeeds, weights are already in model params. # A non-empty dict contains size-mismatched tensors that # should be merged via the standard disk pipeline. weights_preloaded = checkpoint_loader.is_weights_preloaded() @@ -481,10 +481,6 @@ def init_meta_tensor(t: torch.Tensor): checkpoint_loader.post_load_apply( model, weights_preloaded=weights_preloaded) - checkpoint_loader.post_load_publish( - model, - checkpoint_dir=checkpoint_dir, - weights_preloaded=weights_preloaded) for module in model.modules(): if hasattr(module, 'post_load_weights') and not getattr( @@ -498,6 +494,10 @@ def init_meta_tensor(t: torch.Tensor): logger.info("moe_load_balancer finalize model done") torch.cuda.current_stream().synchronize() + checkpoint_loader.post_load_publish( + model, + checkpoint_dir=checkpoint_dir, + weights_preloaded=weights_preloaded) return model, moe_load_balancer diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 0a7b2297e3b5..abf5ec45f413 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3695,12 +3695,12 @@ class LoadFormat(Enum): class ModelExpressConfig(StrictBaseModel): - """Prototype configuration for ModelExpress (MX) weight transfer.""" + """Prototype configuration for ModelExpress weight transfer.""" server_url: Optional[str] = Field( default=None, - description="URL of the MX (ModelExpress) server for P2P weight " - "transfer. When set together with checkpoint_format='MX', enables " + description="URL of the ModelExpress server for P2P weight " + "transfer. When set together with checkpoint_format='modelexpress', enables " "GPU-to-GPU weight transfer from a running source instance, bypassing " "disk I/O. When the server is unreachable, loading falls back to the " "standard HuggingFace checkpoint path.", @@ -3709,18 +3709,18 @@ class ModelExpressConfig(StrictBaseModel): server_query_timeout_s: Optional[NonNegativeInt] = Field( default=None, - description="Timeout in seconds for upstream MxLiveWeightLoader source " - "discovery. When unset, TRT-LLM first probes for existing sources: " - "no source uses a short 30-second fallback cap, while an existing " - "source uses modelexpress's default wait for long donor loads.", + description="Timeout in seconds for ModelExpress TRT-LLM source " + "discovery before falling back to standard checkpoint loading. " + "TRT-LLM passes this value through; ModelExpress owns the probe and " + "fallback behavior.", status="prototype", ) preshard_strategy: str = Field( default="per_module", - description="How to inform TRT-LLM that MX-delivered weights are already " + description="How to inform TRT-LLM that ModelExpress-delivered weights are already " "TP-sharded for the local rank. Only 'per_module' is supported in " - "this MX-only PR; 'global' requires LoadFormat.PRESHARDED.", + "this ModelExpress-only PR; 'global' requires LoadFormat.PRESHARDED.", status="prototype", ) @@ -3988,7 +3988,7 @@ class TorchLlmArgs(BaseLlmArgs): mx_config: ModelExpressConfig = Field( default_factory=ModelExpressConfig, - description="ModelExpress (MX) P2P checkpoint loading config.", + description="ModelExpress P2P checkpoint loading config.", status="prototype", ) @@ -4296,12 +4296,13 @@ def validate_checkpoint_format(self): @model_validator(mode="after") def validate_mx_config(self) -> 'TorchLlmArgs': - # When MX is the active checkpoint format and the user did not + # When ModelExpress is the active checkpoint format and the user did not # explicitly set ``mx_config.server_url``, honor the ``MODEL_EXPRESS_URL`` # env var that the upstream ``modelexpress`` library reads. This - # lets orchestrators configure MX via the environment while keeping + # lets orchestrators configure ModelExpress via the environment while keeping # the resolved value visible on ``llm_args.mx_config.server_url``. - if self.checkpoint_format == "MX" and self.mx_config.server_url is None: + if (self.checkpoint_format == "modelexpress" + and self.mx_config.server_url is None): env_url = os.environ.get("MODEL_EXPRESS_URL") if env_url: logger.info( @@ -4309,11 +4310,12 @@ def validate_mx_config(self) -> 'TorchLlmArgs': "from environment.", env_url) self.mx_config.server_url = env_url - if self.mx_config.server_url is not None and self.checkpoint_format != "MX": + if (self.mx_config.server_url is not None + and self.checkpoint_format != "modelexpress"): logger.warning( "mx_config.server_url is set but checkpoint_format is '%s', not " - "'MX'. The MX config will be ignored. Set " - "checkpoint_format='MX' to enable MX P2P weight transfer.", + "'modelexpress'. The ModelExpress config will be ignored. Set " + "checkpoint_format='modelexpress' to enable P2P weight transfer.", self.checkpoint_format) return self diff --git a/tests/unittest/_torch/models/checkpoints/mx/test_mx_checkpoint_loader.py b/tests/unittest/_torch/models/checkpoints/mx/test_mx_checkpoint_loader.py index 7ffea89b7773..9bd193e84007 100644 --- a/tests/unittest/_torch/models/checkpoints/mx/test_mx_checkpoint_loader.py +++ b/tests/unittest/_torch/models/checkpoints/mx/test_mx_checkpoint_loader.py @@ -1,684 +1,69 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -"""Unit tests for ``MXCheckpointLoader`` (``checkpoint_format='MX'``). -These tests intentionally do NOT exercise the upstream ``modelexpress`` -library. Tests for the import-failure fallback path mock -``modelexpress.*`` symbols out of ``sys.modules`` so the assertion is -about *our* fallback behavior, not about the upstream API. -""" +"""Unit tests for the in-tree ModelExpress checkpoint loader shim.""" -import os import sys -from contextlib import ExitStack -from unittest.mock import MagicMock, patch +import types import pytest -from tensorrt_llm._torch.models.checkpoints.auto_mapper import AutoCheckpointMapper from tensorrt_llm._torch.models.checkpoints.base_checkpoint_loader import BaseCheckpointLoader -from tensorrt_llm._torch.models.checkpoints.hf.checkpoint_loader import HfCheckpointLoader -from tensorrt_llm._torch.models.checkpoints.hf.qwen3_next_weight_mapper import ( - Qwen3NextHfWeightMapper, -) -from tensorrt_llm._torch.models.checkpoints.hf.weight_mapper import HfWeightMapper -from tensorrt_llm._torch.models.checkpoints.mx.checkpoint_loader import ( - MXCheckpointLoader, - _normalize_model_identity, - _resolve_mx_model_name, -) +from tensorrt_llm._torch.models.checkpoints.mx.checkpoint_loader import MXCheckpointLoader -# --------------------------------------------------------------------------- -# Construction & static properties -# --------------------------------------------------------------------------- +class FakeModelexpressMXCheckpointLoader: + def __init__(self, **kwargs): + self.kwargs = kwargs -class TestConstruction: - def test_no_args_constructs(self): - loader = MXCheckpointLoader() - assert loader.mx_server_url is None - assert loader.p2p_succeeded is False - def test_mx_server_url_stored(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - assert loader.mx_server_url == "http://mx:8001" - assert loader.p2p_succeeded is False +def _install_fake_modelexpress_loader(monkeypatch): + fake_modelexpress = types.ModuleType("modelexpress") + fake_engines = types.ModuleType("modelexpress.engines") + fake_trtllm = types.ModuleType("modelexpress.engines.trtllm") + fake_loader = types.ModuleType("modelexpress.engines.trtllm.loader") + fake_loader.MXCheckpointLoader = FakeModelexpressMXCheckpointLoader - def test_query_timeout_stored(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001", query_timeout_s=900) - assert loader.query_timeout_s == 900 + fake_modelexpress.engines = fake_engines + fake_engines.trtllm = fake_trtllm + fake_trtllm.loader = fake_loader - def test_subclasses_hf_loader(self): - # Inheriting from HfCheckpointLoader is what gives us free disk - # fallback. Don't break this. - loader = MXCheckpointLoader() - assert isinstance(loader, HfCheckpointLoader) - assert isinstance(loader, BaseCheckpointLoader) + monkeypatch.setitem(sys.modules, "modelexpress", fake_modelexpress) + monkeypatch.setitem(sys.modules, "modelexpress.engines", fake_engines) + monkeypatch.setitem(sys.modules, "modelexpress.engines.trtllm", fake_trtllm) + monkeypatch.setitem(sys.modules, "modelexpress.engines.trtllm.loader", fake_loader) - def test_checkpoint_format_property(self): - loader = MXCheckpointLoader() - assert loader.checkpoint_format == "MX" - def test_checkpoint_format_backing_attr(self): - # Several call sites read ``self._checkpoint_format`` directly - # (instead of going through the property). The constructor must - # align the backing attribute with the property override. - loader = MXCheckpointLoader() - assert loader._checkpoint_format == "MX" +def test_shim_instantiates_modelexpress_loader(monkeypatch): + _install_fake_modelexpress_loader(monkeypatch) - def test_p2p_succeeded_property_initial(self): - loader = MXCheckpointLoader() - assert loader.p2p_succeeded is False + loader = MXCheckpointLoader(mx_server_url="mx.example:8001") + assert isinstance(loader, FakeModelexpressMXCheckpointLoader) + assert loader.kwargs == {"mx_server_url": "mx.example:8001"} -# --------------------------------------------------------------------------- -# Registry -# --------------------------------------------------------------------------- +def test_registry_instantiates_modelexpress_loader(monkeypatch): + _install_fake_modelexpress_loader(monkeypatch) -class TestRegistry: - def test_registered_under_mx(self): - # ``ModelLoader.load`` resolves a checkpoint loader from - # ``checkpoint_format`` via ``BaseCheckpointLoader.get``. PR #13045 - # registers ``MX`` via ``@register_checkpoint_loader("MX")``. - # The real call site (in _construct_checkpoint_loader) passes - # weight_loader=, weight_mapper=, config_loader=, plus optional - # mx_server_url= — so test with the same call shape. - loader = BaseCheckpointLoader.get( - checkpoint_format="MX", - weight_loader=None, - weight_mapper=None, - config_loader=None, - mx_server_url="http://mx:8001", - query_timeout_s=900, - ) - assert isinstance(loader, MXCheckpointLoader) - assert loader.checkpoint_format == "MX" - assert loader.mx_server_url == "http://mx:8001" - assert loader.query_timeout_s == 900 - - -class TestMxMapperFallback: - def test_arch_specific_mx_mapper_falls_back_to_hf_mapper(self): - mapper = AutoCheckpointMapper.get("MX", "Qwen3NextForCausalLM") - assert isinstance(mapper, Qwen3NextHfWeightMapper) - - def test_unknown_arch_uses_default_mx_mapper(self): - mapper = AutoCheckpointMapper.get("MX", "UnknownArchitecture") - assert isinstance(mapper, HfWeightMapper) - - -# --------------------------------------------------------------------------- -# load_weights — disk-fallback paths (no upstream library involved) -# --------------------------------------------------------------------------- - - -class TestLoadWeightsFallback: - """Disk-fallback paths that should not touch the upstream MX library. - - All four fallback triggers share the same observable contract: - ``p2p_succeeded`` stays False, the parent ``HfCheckpointLoader. - load_weights`` is invoked exactly once, and its return value is - propagated unchanged. We parameterize the trigger to keep that - contract in one place. - """ - - # Trigger setup builders. - @staticmethod - def _no_url(stack): # noqa: ARG004 — stack unused for this trigger - return MXCheckpointLoader(), {"model": MagicMock()} - - @staticmethod - def _no_model(stack): # noqa: ARG004 - return MXCheckpointLoader(mx_server_url="http://mx:8001"), {} - - @staticmethod - def _modelexpress_unavailable(stack): - stack.enter_context(_block_modelexpress()) - return (MXCheckpointLoader(mx_server_url="http://mx:8001"), {"model": MagicMock()}) - - @staticmethod - def _upstream_raises(stack): - fake_mx = _build_fake_modelexpress(load_weights_side_effect=RuntimeError("boom")) - stack.enter_context(_install_fake_modelexpress(fake_mx)) - return (MXCheckpointLoader(mx_server_url="http://mx:8001"), {"model": MagicMock()}) - - @pytest.mark.parametrize( - "trigger_id, setup", - [ - ("no_mx_server_url", _no_url), - ("no_model_kwarg", _no_model), - ("modelexpress_not_installed", _modelexpress_unavailable), - ("upstream_raises", _upstream_raises), - ], - ) - def test_falls_back_to_disk(self, trigger_id, setup): - sentinel = {"disk-load": "result"} - with ExitStack() as stack: - loader, extra_kwargs = setup(stack) - mock_super_load = stack.enter_context( - patch.object(HfCheckpointLoader, "load_weights", return_value=sentinel) - ) - - result = loader.load_weights("/nonexistent", mapping=MagicMock(), **extra_kwargs) - - assert result is sentinel, ( - f"trigger={trigger_id}: production code must propagate the " - "parent loader's return value unchanged" - ) - assert loader.p2p_succeeded is False, ( - f"trigger={trigger_id}: p2p_succeeded must stay False on any fallback path" - ) - mock_super_load.assert_called_once() - - -# --------------------------------------------------------------------------- -# load_weights — MX-success and mixed-success paths (mocked upstream) -# --------------------------------------------------------------------------- - - -class TestLoadWeightsMxPath: - def test_p2p_full_success_returns_empty_dict(self): - # Empty fallback dict means MX delivered all weights into model - # params; ``ModelLoader`` interprets the empty dict + p2p_succeeded - # flag as "skip the standard weight-mapping pipeline". - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress(load_weights_return={}) - mapping = MagicMock(name="mapping") - model = MagicMock(name="model") - - with _install_fake_modelexpress(fake_mx): - result = loader.load_weights("/nonexistent", mapping=mapping, model=model) - - assert result == {} - assert loader.p2p_succeeded is True - - # Verify the integration contract with the upstream library: - # 1. Constructed MxLiveWeightLoader with our mx_server_url. - fake_mx.trtllm_live_transfer.MxLiveWeightLoader.assert_called_once_with( - mx_server="http://mx:8001" - ) - # 2. Called load_weights with the right positional/keyword args. - weight_loader_instance = fake_mx.trtllm_live_transfer.MxLiveWeightLoader.return_value - weight_loader_instance.load_weights.assert_called_once_with( - "/nonexistent", mapping=mapping, model=model - ) - - def test_mixed_success_returns_fallback_weights(self): - # When MX returns a non-empty fallback dict (size-mismatched - # tensors), keep the P2P transfer and let ModelLoader merge these - # tensors through the standard disk pipeline. - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fallback = {"some.weight": MagicMock()} - fake_mx = _build_fake_modelexpress(load_weights_return=fallback) - - with ( - _install_fake_modelexpress(fake_mx), - patch.object(HfCheckpointLoader, "load_weights") as mock_super_load, - ): - result = loader.load_weights("/nonexistent", mapping=MagicMock(), model=MagicMock()) - - assert loader.p2p_succeeded is True - assert result is fallback - mock_super_load.assert_not_called() - - -# --------------------------------------------------------------------------- -# publish_as_source — env-var dance and graceful no-op -# --------------------------------------------------------------------------- - - -class TestPublishAsSource: - def test_no_mx_server_url_is_noop(self): - loader = MXCheckpointLoader() # mx_server_url is None - # Any attempt to import modelexpress would raise here, so we - # don't even need to mock. - with _block_modelexpress(): - loader.publish_as_source(MagicMock()) # must not raise - - def test_modelexpress_unavailable_is_noop(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - with _block_modelexpress(): - loader.publish_as_source(MagicMock()) # must not raise - - def test_publish_called_with_model(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress() - model = MagicMock(name="model") - - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(model) - - fake_mx.trtllm_live_transfer.publish_model_params.assert_called_once_with(model) - - def test_env_var_set_during_publish_then_restored(self): - loader = MXCheckpointLoader(mx_server_url="http://mx-instance:9999") - captured_env = {} - - def _capture(model): - captured_env["MODEL_EXPRESS_URL"] = os.environ.get("MODEL_EXPRESS_URL") - - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - - prior = os.environ.pop("MODEL_EXPRESS_URL", None) - try: - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) - finally: - if prior is not None: - os.environ["MODEL_EXPRESS_URL"] = prior - - # During the call, our per-loader URL was visible. - assert captured_env["MODEL_EXPRESS_URL"] == "http://mx-instance:9999" - # After the call, the env var is back to the pre-call state. - assert "MODEL_EXPRESS_URL" not in os.environ - - def test_env_var_restored_to_prior_value(self): - loader = MXCheckpointLoader(mx_server_url="http://mx-instance:9999") - fake_mx = _build_fake_modelexpress() - prior = os.environ.get("MODEL_EXPRESS_URL") - os.environ["MODEL_EXPRESS_URL"] = "http://prior-value:1234" - try: - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) - assert os.environ["MODEL_EXPRESS_URL"] == "http://prior-value:1234" - finally: - if prior is None: - os.environ.pop("MODEL_EXPRESS_URL", None) - else: - os.environ["MODEL_EXPRESS_URL"] = prior - - def test_publish_exception_swallowed(self): - # publish_as_source is a best-effort hook; an upstream exception - # must NOT take down model loading. - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress(publish_side_effect=RuntimeError("upstream went away")) - - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) # must not raise - - -# --------------------------------------------------------------------------- -# Helpers — fake modelexpress modules and import blockers -# --------------------------------------------------------------------------- - - -def _modelexpress_module_names(): - return [ - "modelexpress", - "modelexpress.trtllm_live_transfer", - ] - - -def _block_modelexpress(): - """Context manager that makes ``import modelexpress`` raise ImportError.""" - saved = {name: sys.modules.get(name) for name in _modelexpress_module_names()} - - class _Blocker: - def __enter__(self): - for name in _modelexpress_module_names(): - # Setting to None makes ``import name`` raise - # ImportError per PEP 328 / sys.modules semantics. - sys.modules[name] = None - return self - - def __exit__(self, exc_type, exc, tb): - for name, prior in saved.items(): - if prior is None: - sys.modules.pop(name, None) - else: - sys.modules[name] = prior - - return _Blocker() - - -def _build_fake_modelexpress( - *, - load_weights_return=None, - load_weights_side_effect=None, - publish_side_effect=None, - source_instances=None, -): - """Build a fake modelexpress module tree mimicking the symbols we use.""" - fake_pkg = MagicMock(name="modelexpress") - - fake_trtllm_live = MagicMock(name="modelexpress.trtllm_live_transfer") - - # MxLiveWeightLoader(mx_server=url).load_weights(ckpt_dir, mapping=, model=) - weight_loader_instance = MagicMock(name="MxLiveWeightLoader instance") - if load_weights_side_effect is not None: - weight_loader_instance.load_weights.side_effect = load_weights_side_effect - else: - weight_loader_instance.load_weights.return_value = ( - load_weights_return if load_weights_return is not None else {} - ) - fake_trtllm_live.MxLiveWeightLoader = MagicMock(return_value=weight_loader_instance) - client_instance = MagicMock(name="MxClient instance") - client_instance.list_sources.return_value = MagicMock(instances=source_instances or []) - fake_trtllm_live.MxClient = MagicMock(return_value=client_instance) - fake_trtllm_live._build_trtllm_identity = MagicMock(return_value=MagicMock()) - - # publish_model_params(model) - if publish_side_effect is not None: - fake_trtllm_live.publish_model_params = MagicMock(side_effect=publish_side_effect) - else: - fake_trtllm_live.publish_model_params = MagicMock() - - fake_pkg.trtllm_live_transfer = fake_trtllm_live - return fake_pkg - - -def _install_fake_modelexpress(fake_pkg): - """Context manager that installs a fake ``modelexpress`` into sys.modules.""" - - class _Installer: - _saved = {} - - def __enter__(self): - for name in _modelexpress_module_names(): - self._saved[name] = sys.modules.get(name) - sys.modules["modelexpress"] = fake_pkg - sys.modules["modelexpress.trtllm_live_transfer"] = fake_pkg.trtllm_live_transfer - return self - - def __exit__(self, exc_type, exc, tb): - for name, prior in self._saved.items(): - if prior is None: - sys.modules.pop(name, None) - else: - sys.modules[name] = prior - - return _Installer() - - -# --------------------------------------------------------------------------- -# Item 2: defensive MX_SOURCE_QUERY_TIMEOUT default -# --------------------------------------------------------------------------- - - -class TestMxSourceQueryTimeoutDefault: - """``load_weights`` caps upstream's source-query timeout on P2P attempts. - - The first replica on a cold cluster should not block for the upstream - default of 1 hour. ``setdefault`` semantics — never overrides an - explicit user value. - """ - - @pytest.fixture(autouse=True) - def _isolated_env(self, monkeypatch): - monkeypatch.delenv("MX_SOURCE_QUERY_TIMEOUT", raising=False) - yield - - def test_no_registered_source_gets_short_default_during_load(self): - def _assert_timeout(*args, **kwargs): - assert os.environ.get("MX_SOURCE_QUERY_TIMEOUT") == "30" - return {} - - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress(load_weights_side_effect=_assert_timeout) - with _install_fake_modelexpress(fake_mx): - loader.load_weights("/nonexistent", mapping=MagicMock(), model=MagicMock()) - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - - def test_existing_source_keeps_upstream_default_when_unset(self): - def _assert_no_timeout(*args, **kwargs): - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - return {} - - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress( - load_weights_side_effect=_assert_no_timeout, - source_instances=[MagicMock()], - ) - with _install_fake_modelexpress(fake_mx): - loader.load_weights("/nonexistent", mapping=MagicMock(), model=MagicMock()) - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - - def test_env_value_preserved(self, monkeypatch): - # If the user/orchestrator already set a value, our defensive - # default must not stomp it. - monkeypatch.setenv("MX_SOURCE_QUERY_TIMEOUT", "120") - - def _assert_env_timeout(*args, **kwargs): - assert os.environ.get("MX_SOURCE_QUERY_TIMEOUT") == "120" - return {} - - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress(load_weights_side_effect=_assert_env_timeout) - with _install_fake_modelexpress(fake_mx): - loader.load_weights("/nonexistent", mapping=MagicMock(), model=MagicMock()) - assert os.environ.get("MX_SOURCE_QUERY_TIMEOUT") == "120" - - def test_configured_timeout_applies_during_load_and_restores_env(self): - def _assert_config_timeout(*args, **kwargs): - assert os.environ.get("MX_SOURCE_QUERY_TIMEOUT") == "900" - return {} - - loader = MXCheckpointLoader(mx_server_url="http://mx:8001", query_timeout_s=900) - fake_mx = _build_fake_modelexpress(load_weights_side_effect=_assert_config_timeout) - with _install_fake_modelexpress(fake_mx): - loader.load_weights("/nonexistent", mapping=MagicMock(), model=MagicMock()) - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - - def test_no_mx_url_does_not_touch_env(self): - # HF-only loads must not surprise users by setting MX-namespaced - # env vars they didn't ask for. - MXCheckpointLoader() # mx_server_url=None - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - - -# --------------------------------------------------------------------------- -# Item 3: model_name plumbing + resolver + publish-side env handoff -# --------------------------------------------------------------------------- - - -class TestModelNameConstructor: - def test_default_model_name_none(self): - loader = MXCheckpointLoader() - assert loader.model_name is None - - def test_model_name_stored_as_string(self): - loader = MXCheckpointLoader(model_name="Qwen/Qwen2.5-72B-Instruct") - assert loader.model_name == "Qwen/Qwen2.5-72B-Instruct" - - def test_model_name_path_coerced_to_string(self, tmp_path): - # Constructor accepts Path (e.g. llm_args.model resolved as Path) - # and stores it as a string for downstream env-var publishing. - path = tmp_path / "my-checkpoint" - loader = MXCheckpointLoader(model_name=path) - assert loader.model_name == str(path) - assert isinstance(loader.model_name, str) - - -class TestNormalizeModelIdentity: - """``_normalize_model_identity`` is the path-vs-id heuristic used by - the resolver. Pinning it down with parametrized cases.""" - - @pytest.mark.parametrize( - "label, value, expected", - [ - # Hub IDs and bare names pass through unchanged. - ("bare_name", "llama-3-70b", "llama-3-70b"), - ("hub_id", "Qwen/Qwen2.5-72B-Instruct", "Qwen/Qwen2.5-72B-Instruct"), - ("nested_hub_id", "meta-llama/Llama-3-70B", "meta-llama/Llama-3-70B"), - # Absolute paths get reduced to basenames. - ("abs_path_simple", "/scratch/local-model", "local-model"), - ("abs_path_nested", "/cache/foo/bar/baz", "baz"), - # Relative-looking paths are treated as paths. - ("dot_relative", "./local-model", "local-model"), - ("home_expansion", "~/models/my-model", "my-model"), - # Empty / sentinel. - ("empty", "", "unknown"), - ], + loader = BaseCheckpointLoader.get( + checkpoint_format="modelexpress", + mx_server_url="mx.example:8001", + model_name="Qwen/Qwen2.5-7B", ) - def test_basic_cases(self, label, value, expected): - assert _normalize_model_identity(value) == expected, f"case={label}" - - def test_hf_snapshot_unmangling(self): - # HF cache layout: ".../models----/snapshots//" - # → "/" (not the commit sha). - snapshot = ( - "/cache/huggingface/hub/models--Qwen--Qwen2.5-72B-Instruct/snapshots/abc123def456789" - ) - assert _normalize_model_identity(snapshot) == "Qwen/Qwen2.5-72B-Instruct" - - def test_hf_snapshot_unmangling_nested_org(self): - # Multi-component HF org names use "--" as the separator. - snapshot = "/cache/hub/models--meta-llama--Llama-3-70B-Instruct/snapshots/sha" - assert _normalize_model_identity(snapshot) == "meta-llama/Llama-3-70B-Instruct" - - -class TestResolveMxModelName: - """``_resolve_mx_model_name`` is the priority-ordered lookup used by - ``publish_as_source``. Verifying the ordering.""" - - @pytest.fixture(autouse=True) - def _isolated_env(self, monkeypatch): - monkeypatch.delenv("MODEL_NAME", raising=False) - yield - - def test_explicit_arg_wins(self, monkeypatch): - monkeypatch.setenv("MODEL_NAME", "from-env") - assert _resolve_mx_model_name("explicit", "/cache/snapshot/abc") == "explicit" - - def test_env_used_when_arg_none(self, monkeypatch): - monkeypatch.setenv("MODEL_NAME", "from-env") - assert _resolve_mx_model_name(None, "/cache/snapshot/abc") == "from-env" - - def test_basename_fallback_when_arg_and_env_missing(self): - assert _resolve_mx_model_name(None, "/scratch/local-model") == "local-model" - - def test_snapshot_fallback_when_arg_and_env_missing(self): - snapshot = "/cache/huggingface/hub/models--Qwen--Qwen2.5-72B-Instruct/snapshots/abc123" - assert _resolve_mx_model_name(None, snapshot) == "Qwen/Qwen2.5-72B-Instruct" - - def test_unknown_when_all_missing(self): - assert _resolve_mx_model_name(None, None) == "unknown" - assert _resolve_mx_model_name("", None) == "unknown" - - def test_explicit_arg_normalized_too(self): - # If the explicit arg looks like a path (e.g. llm_args.model - # was a Path), it gets normalized too. - assert _resolve_mx_model_name("/scratch/explicit-path", "/cache/ignored") == "explicit-path" - - -class TestPublishAsSourceModelName: - """``publish_as_source`` must set ``MODEL_NAME`` from the resolved - identity so upstream's ``publish_model_params`` reads it via env, - and must restore the prior env value afterwards. - """ - - @pytest.fixture(autouse=True) - def _isolated_env(self, monkeypatch): - monkeypatch.delenv("MODEL_NAME", raising=False) - monkeypatch.delenv("MODEL_EXPRESS_URL", raising=False) - monkeypatch.delenv("MX_SOURCE_QUERY_TIMEOUT", raising=False) - yield - - def test_uses_explicit_constructor_model_name(self): - loader = MXCheckpointLoader( - mx_server_url="http://mx:8001", - model_name="Qwen/Qwen2.5-72B-Instruct", - ) - captured = {} - - def _capture(model): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") - captured["MODEL_EXPRESS_URL"] = os.environ.get("MODEL_EXPRESS_URL") - - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) - - assert captured["MODEL_NAME"] == "Qwen/Qwen2.5-72B-Instruct" - assert captured["MODEL_EXPRESS_URL"] == "http://mx:8001" - # Both env vars restored to the (unset) prior state. - assert "MODEL_NAME" not in os.environ - assert "MODEL_EXPRESS_URL" not in os.environ - - def test_falls_back_to_checkpoint_dir_basename(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - # No constructor model_name, no MODEL_NAME env → use basename. - captured = {} - - def _capture(model): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") - - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock(), checkpoint_dir="/scratch/local-model") - - assert captured["MODEL_NAME"] == "local-model" - - def test_unmangles_hf_snapshot_path(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - snapshot = ( - "/cache/huggingface/hub/models--Qwen--Qwen2.5-72B-Instruct/snapshots/abc123def456" - ) - captured = {} - - def _capture(model): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") - - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock(), checkpoint_dir=snapshot) - - # Critical: NOT the commit hash, the human-readable Hub-ID form. - assert captured["MODEL_NAME"] == "Qwen/Qwen2.5-72B-Instruct" - assert captured["MODEL_NAME"] != "abc123def456" - - def test_constructor_model_name_takes_priority_over_env(self, monkeypatch): - # Explicit constructor value > existing MODEL_NAME env. - monkeypatch.setenv("MODEL_NAME", "from-env") - loader = MXCheckpointLoader( - mx_server_url="http://mx:8001", - model_name="explicit", - ) - captured = {} - - def _capture(model): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") - - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) - - assert captured["MODEL_NAME"] == "explicit" - # Restored to the prior env value, not unset. - assert os.environ.get("MODEL_NAME") == "from-env" - - def test_env_used_when_no_constructor_value(self, monkeypatch): - monkeypatch.setenv("MODEL_NAME", "from-env-only") - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - captured = {} - - def _capture(model): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") - - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) - - assert captured["MODEL_NAME"] == "from-env-only" - assert os.environ.get("MODEL_NAME") == "from-env-only" - def test_unknown_when_all_sources_missing(self): - # No constructor, no env, no checkpoint_dir → upstream sentinel. - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - captured = {} + assert isinstance(loader, FakeModelexpressMXCheckpointLoader) + assert loader.kwargs == { + "mx_server_url": "mx.example:8001", + "model_name": "Qwen/Qwen2.5-7B", + } - def _capture(model): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) +def test_shim_requires_modelexpress(monkeypatch): + monkeypatch.setitem(sys.modules, "modelexpress", None) + monkeypatch.setitem(sys.modules, "modelexpress.engines", None) + monkeypatch.setitem(sys.modules, "modelexpress.engines.trtllm", None) + monkeypatch.setitem(sys.modules, "modelexpress.engines.trtllm.loader", None) - assert captured["MODEL_NAME"] == "unknown" + with pytest.raises(ImportError, match="requires the modelexpress Python package"): + MXCheckpointLoader() diff --git a/tests/unittest/_torch/pyexecutor/test_model_loader_mx.py b/tests/unittest/_torch/pyexecutor/test_model_loader_mx.py index ae790dd04255..6de7858be710 100644 --- a/tests/unittest/_torch/pyexecutor/test_model_loader_mx.py +++ b/tests/unittest/_torch/pyexecutor/test_model_loader_mx.py @@ -1,6 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Unit tests for MX-specific branches in ``ModelLoader``.""" +"""Unit tests for ModelExpress-specific branches in ``ModelLoader``.""" from contextlib import contextmanager, nullcontext from types import SimpleNamespace @@ -36,7 +36,7 @@ def __init__(self, events): self._events = events def _apply(self, fn): - # The test is about ModelLoader's MX branching, not CUDA allocation. + # The test is about ModelLoader's ModelExpress branching, not CUDA allocation. return self def to(self, *args, **kwargs): @@ -84,7 +84,7 @@ def _make_loader(monkeypatch, *, events, spec_config=None): monkeypatch.setattr( torch.cuda, "current_stream", - lambda: SimpleNamespace(synchronize=lambda: None), + lambda: SimpleNamespace(synchronize=lambda: events.append("sync")), ) return loader @@ -93,9 +93,12 @@ def test_mx_success_initializes_mapper_skips_weight_mapping_and_reload_works(mon events = [] loader = _make_loader(monkeypatch, events=events) checkpoint_loader = MagicMock(name="checkpoint_loader") - checkpoint_loader.checkpoint_format = "MX" + checkpoint_loader.checkpoint_format = "modelexpress" checkpoint_loader.is_weights_preloaded.return_value = True checkpoint_loader.load_weights.return_value = {} + checkpoint_loader.post_load_publish.side_effect = lambda *args, **kwargs: events.append( + "post_load_publish" + ) model, _ = loader.load("/ckpt", checkpoint_loader) @@ -109,8 +112,9 @@ def test_mx_success_initializes_mapper_skips_weight_mapping_and_reload_works(mon checkpoint_loader.post_load_publish.assert_called_once_with( model, checkpoint_dir="/ckpt", weights_preloaded=True ) + assert events[-2:] == ["sync", "post_load_publish"] - # reload() uses self.weight_mapper unconditionally; MX success must + # reload() uses self.weight_mapper unconditionally; ModelExpress success must # initialize it even though the initial load skipped _call_load_weights. loader.reload(model, {"reloaded": MagicMock()}) assert loader._call_load_weights.call_count == 1 @@ -120,10 +124,13 @@ def test_mx_partial_fallback_merges_returned_weights(monkeypatch): events = [] loader = _make_loader(monkeypatch, events=events) checkpoint_loader = MagicMock(name="checkpoint_loader") - checkpoint_loader.checkpoint_format = "MX" + checkpoint_loader.checkpoint_format = "modelexpress" checkpoint_loader.is_weights_preloaded.return_value = True fallback_weights = {"mismatched.weight": MagicMock()} checkpoint_loader.load_weights.return_value = fallback_weights + checkpoint_loader.post_load_publish.side_effect = lambda *args, **kwargs: events.append( + "post_load_publish" + ) model, _ = loader.load("/ckpt", checkpoint_loader) @@ -135,16 +142,20 @@ def test_mx_partial_fallback_merges_returned_weights(monkeypatch): checkpoint_loader.post_load_publish.assert_called_once_with( model, checkpoint_dir="/ckpt", weights_preloaded=True ) + assert events[-2:] == ["sync", "post_load_publish"] def test_mx_fallback_runs_standard_weight_mapping(monkeypatch): events = [] loader = _make_loader(monkeypatch, events=events) checkpoint_loader = MagicMock(name="checkpoint_loader") - checkpoint_loader.checkpoint_format = "MX" + checkpoint_loader.checkpoint_format = "modelexpress" checkpoint_loader.is_weights_preloaded.return_value = False checkpoint_loader.load_weights.return_value = {"weight": MagicMock()} checkpoint_loader.get_initialized_weight_mapper.return_value = MagicMock() + checkpoint_loader.post_load_publish.side_effect = lambda *args, **kwargs: events.append( + "post_load_publish" + ) model, _ = loader.load("/ckpt", checkpoint_loader) @@ -154,3 +165,4 @@ def test_mx_fallback_runs_standard_weight_mapping(monkeypatch): checkpoint_loader.post_load_publish.assert_called_once_with( model, checkpoint_dir="/ckpt", weights_preloaded=False ) + assert events[-2:] == ["sync", "post_load_publish"]