diff --git a/docs/source/features/model-express.md b/docs/source/features/model-express.md index b0b832cf42b6..9360d24fe531 100644 --- a/docs/source/features/model-express.md +++ b/docs/source/features/model-express.md @@ -120,7 +120,7 @@ uses a metadata-only view of the donor's canonical snapshot and contains no weight shards. A positive result therefore requires direct transfer; disk fallback cannot accidentally satisfy the test. -Run the TP=1 smoke test against an isolated ModelExpress 0.4.1 service with +Run the TP=1 smoke test against an isolated ModelExpress 0.5.1 service with NIXL enabled: ```bash @@ -144,7 +144,7 @@ row. `TRTLLM_MX_E2E_TIMEOUT_S` controls the 1200-second timeout used for the baseline worker, receiver worker, and donor-readiness wait; increase it for slow model storage or startup. -The dedicated H100 CI stages own isolated Redis and ModelExpress 0.4.1 +The dedicated H100 CI stages own isolated Redis and ModelExpress 0.5.1 sidecars. The two-GPU TP=1 stage is classified as multi-GPU: it runs automatically in post-merge pipelines or when a multi-GPU file changes, while direct pre-merge dispatch requires the `ci: full pre-merge approved` label. @@ -190,11 +190,14 @@ When adding an ABI ID: ## Installation -The official TensorRT LLM release container includes the MX Python client. No -additional Python package installation is required in that container. MX -remains opt-in at runtime: TensorRT LLM uses the client only when the MX -checkpoint-loading path and a server URL are configured. Installing the client -does not expand the model support scope described above. +TensorRT LLM release containers that include this feature already install a +compatible MX Python client; no additional client installation is needed for +P2P transfer. For earlier TensorRT LLM releases, use a release or container +built with this feature; upgrading the MX package alone does not add the +missing TensorRT LLM integration. MX remains opt-in at runtime: TensorRT LLM +uses the client only when the MX checkpoint-loading path and a server URL are +configured. Installing the client does not expand the model support scope +described above. For pip installations outside the official release container, install the MX Python client through the optional `mx` extra: @@ -203,15 +206,14 @@ Python client through the optional `mx` extra: pip install "tensorrt-llm[mx]" ``` -The extra accepts ModelExpress client versions `>=0.4.1,<0.6.0`. Version -`0.4.1` is the minimum client API qualified by this integration, while the -upper bound prevents resolving unqualified `0.6.0` or newer client APIs. -Deploy a compatible MX server version. +The extra accepts ModelExpress client versions `>=0.5.1,<0.6.0`. Version +`0.5.1` is the minimum client release that provides the TensorRT LLM adapter, +while the upper bound prevents resolving unqualified `0.6.0` or newer client +APIs. Deploy a compatible MX server version. The extra can be added to an existing TensorRT LLM installation. If the MX loading path is configured but the client cannot be imported, TensorRT LLM -fails with an actionable installation message instead of silently loading from -the Hugging Face checkpoint. Source discovery and transfer failures continue to -use the Hugging Face fallback described above. +logs a warning and uses the Hugging Face fallback described above. Source +discovery and transfer failures use the same fallback. ## Deploy the MX Service @@ -236,7 +238,7 @@ docker run -d --name modelexpress-server \ -e MODEL_EXPRESS_LOG_LEVEL=info \ -e MX_METADATA_BACKEND=redis \ -e REDIS_URL=redis://modelexpress-redis:6379 \ - nvcr.io/nvidia/ai-dynamo/modelexpress-server:0.4.1 + nvcr.io/nvidia/ai-dynamo/modelexpress-server:0.5.1 ``` ## Configure TensorRT LLM @@ -272,7 +274,7 @@ path. | Field | Default | Description | |-------|---------|-------------| | `mx_config.server_url` | `null` | URL of the separately managed MX server. | -| `mx_config.server_query_timeout_s` | `null` | Timeout for MX source discovery. When unset, TensorRT LLM uses a short fallback cap when no source exists and otherwise lets MX wait for long donor loads. | +| `mx_config.server_query_timeout_s` | `null` | Deprecated and ignored. MX checks once for a compatible source, then falls back to native checkpoint loading. | ## Notes and Limitations diff --git a/jenkins/L0_Test.groovy b/jenkins/L0_Test.groovy index bf91395b28da..036d2732c210 100644 --- a/jenkins/L0_Test.groovy +++ b/jenkins/L0_Test.groovy @@ -73,7 +73,7 @@ ARTIFACTORY_CREDENTIALS_ID = "trtllm-artifactory-credentials" // DLFW torch image DLFW_IMAGE = "urm.nvidia.com/docker/nvidia/pytorch:26.05-py3" -MODEL_EXPRESS_VERSION = "0.4.1" +MODEL_EXPRESS_VERSION = "0.5.1" MODEL_EXPRESS_NIXL_VERSION = "1.4.0" MODEL_EXPRESS_SERVER_IMAGE = "urm.nvidia.com/docker/nvidia/ai-dynamo/modelexpress-server:${MODEL_EXPRESS_VERSION}" MODEL_EXPRESS_REDIS_IMAGE = "urm.nvidia.com/docker/redis:7-alpine" @@ -3704,7 +3704,7 @@ def createKubernetesPodConfig(image, type, arch = "amd64", gpuCount = 1, perfMod - name: TRTLLM_MX_E2E_REQUIRED value: "1" """ - // Mirrors the ModelExpress v0.4.1 Redis deployment and image contract. + // Mirrors the ModelExpress Redis deployment and image contract. // The image exposes /app/modelexpress-server and accepts the port/backend settings below. // Use regular containers because the Jenkins Kubernetes launcher does not // reliably attach to pods containing restartable init-container sidecars. @@ -4993,7 +4993,7 @@ def runLLMTestlistOnPlatformImpl(pipeline, platform, testList, config=VANILLA_CO } if (stageName.contains("-ModelExpress-")) { trtllm_utils.llmExecStepWithRetry(pipeline, script: "pip3 install modelexpress==${MODEL_EXPRESS_VERSION}") - // ModelExpress 0.4.1 imports nixl._api, while requirements-dev.txt + // ModelExpress imports nixl._api, while requirements-dev.txt // installs only the nixl-cu13 backend. Install the matching // namespace shim without pulling the unused CUDA 12 backend. trtllm_utils.llmExecStepWithRetry(pipeline, script: "pip3 install --no-deps nixl==${MODEL_EXPRESS_NIXL_VERSION}") diff --git a/setup.py b/setup.py index 205735a4daae..702b6729a9b7 100644 --- a/setup.py +++ b/setup.py @@ -148,7 +148,7 @@ def has_ext_modules(self): Path("requirements-dev-windows.txt" if on_windows else "requirements-dev.txt")) openengine_deps, _ = parse_requirements(Path("requirements-openengine.txt")) -mx_deps = ["modelexpress>=0.4.1,<0.6.0"] +mx_deps = ["modelexpress>=0.5.1,<0.6.0"] # Gateway protocol adapters are opt-in extras: the default installation must # not carry any gateway protobuf package. Each gateway owns a dedicated # requirements-.txt as the single source of truth for its pins; CI diff --git a/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py index 714396bd9142..6f83b1e75264 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mx/checkpoint_loader.py @@ -12,186 +12,31 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""MX (ModelExpress) checkpoint loader. +"""ModelExpress 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. This class only calls them at the right -points in TRT-LLM's loading lifecycle. - -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. +TensorRT-LLM owns model construction, compatibility identity, post-transform +qualification, and post-load lifecycle. ModelExpress owns source selection, +RDMA transfer, native fallback selection, publication, and transport cleanup +through its shared load-strategy chain. """ -import inspect -import json import logging import os -import threading -import traceback -from contextlib import contextmanager -from enum import Enum -from pathlib import Path -from typing import Any, Callable, Iterator, MutableMapping, Optional, Protocol, Type, Union +from importlib import import_module +from typing import Any, Optional 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._torch.weight_sharing import ( - IdentityCheckPolicy, - SourceIdentity, - check_weight_sharing_compatibility, -) +from tensorrt_llm._torch.weight_sharing import SOURCE_IDENTITY_FORMAT_VERSION, SourceIdentity 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 the per-loader `query_timeout_s` setting. -# Tracked as MX-4 in §15 (non-blocking source-query API upstream). -_MX_SOURCE_QUERY_TIMEOUT_DEFAULT_S = "30" -# ModelExpress 0.4.1 reads transfer configuration from process-wide -# environment variables and exposes a module-level identity builder. Keep all -# temporary mutation of that shared state in one critical section. -_MX_TRANSFER_STATE_LOCK = threading.Lock() -_MX_SOURCE_IDENTITY_METADATA_KEY = "trtllm_source_identity" -_MX_WEIGHT_LAYOUT_METADATA_KEY = "trtllm_weight_layout" -_MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY = "trtllm_transform_protocol_version" -_MX_TRANSFORM_ABI_ID_METADATA_KEY = "trtllm_transform_abi_id" -_MX_WEIGHT_LAYOUT_POST_TRANSFORM = "post_transform" -_MX_STAGED_TRANSFORM_PROTOCOL_VERSION = 1 - - -class _MxWeightLayoutStatus(Enum): - PRE_TRANSFORM = "pre_transform" - POST_TRANSFORM_SUPPORTED = "post_transform_supported" - UNSUPPORTED = "unsupported" - - -class _MxSourceIdentity(Protocol): - """Subset of ModelExpress's protobuf SourceIdentity used by this adapter.""" - - extra_parameters: MutableMapping[str, str] - - -@contextmanager -def _temporary_env(key: str, value: Optional[str]) -> Iterator[None]: - """Temporarily set one environment variable when a value is provided.""" - 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 - - -def _serialize_source_identity(identity: SourceIdentity) -> str: - """Serialize TRT-LLM's layout identity for MX's identity map.""" - payload = identity.to_dict() - # `model_name` is a cleartext discovery descriptor and is deliberately - # excluded from SourceIdentity compatibility checks. The outer MX identity - # already carries the normalized model name; embedding a local checkpoint - # path here would make otherwise-compatible no-shards receivers hash to a - # different MX source. - payload.pop("model_name", None) - return json.dumps( - payload, - sort_keys=True, - separators=(",", ":"), - ) - - -def _attach_trtllm_metadata_to_mx_identity( - mx_identity: _MxSourceIdentity, source_identity: Optional[SourceIdentity] -) -> _MxSourceIdentity: - """Attach TRT-LLM compatibility metadata to an MX SourceIdentity.""" - if source_identity is None: - return mx_identity - - extra_parameters = getattr(mx_identity, "extra_parameters", None) - if extra_parameters is None: - raise RuntimeError( - "MX SourceIdentity has no extra_parameters field; cannot attach " - "TRT-LLM SourceIdentity for compatibility filtering." - ) - - try: - for key, value in _build_mx_source_metadata(source_identity).items(): - extra_parameters[key] = value - except (AttributeError, TypeError, ValueError) as e: - raise RuntimeError( - "Failed to attach TRT-LLM compatibility metadata to MX " - "SourceIdentity; MX P2P compatibility filtering will reject " - "this source." - ) from e - return mx_identity - - -@contextmanager -def _patched_trtllm_identity_builder( - mx_transfer: Any, source_identity: Optional[SourceIdentity] -) -> Iterator[None]: - """Temporarily wrap upstream TRT-LLM identity construction.""" - original = getattr(mx_transfer, "_build_trtllm_identity", None) - if source_identity is None or not callable(original): - yield - return - - def _wrapped_build_identity(*args: Any, **kwargs: Any) -> _MxSourceIdentity: - return _attach_trtllm_metadata_to_mx_identity( - original(*args, **kwargs), - source_identity, - ) - - mx_transfer._build_trtllm_identity = _wrapped_build_identity - try: - yield - finally: - mx_transfer._build_trtllm_identity = original - - -def _close_mx_client(client: Any) -> None: - """Close a best-effort MX discovery client without masking its result.""" - if client is None: - return - close = getattr(client, "close", None) - if not callable(close): - return - try: - close() - except Exception: - logger.warning( - f"Failed to close MX discovery client; continuing with the " - f"completed probe result.\n{traceback.format_exc()}" - ) - - -def _synchronize_cuda_for_mx_publish() -> None: - """Finish pending CUDA writes before exposing source buffers through MX.""" - import torch - - if torch.cuda.is_initialized(): - torch.cuda.synchronize() - def _enable_mx_transfer_logging() -> None: - """Enable upstream INFO records when per-rank transfer logs are requested.""" + """Enable ModelExpress INFO records for requested per-rank transfer logs.""" if not os.environ.get("MX_TRANSFER_LOG_DIR"): return @@ -202,24 +47,7 @@ def _enable_mx_transfer_logging() -> None: @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 after `post_load_weights()` runs, together with - metadata that lets compatible targets skip one-shot post-load transforms. - - When the MX server is unavailable, this loader transparently falls back - to standard HuggingFace checkpoint loading via the parent - `HfCheckpointLoader`. A missing MX client is treated as a configuration - error and reported with an actionable installation command. - - 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. - """ + """Load a TRT-LLM shard through ModelExpress, with native HF fallback.""" def __init__( self, @@ -228,90 +56,33 @@ def __init__( 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). Transfer and publish paths resolve 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 self._post_transform_weights_preloaded = False self._source_identity_compatible_for_last_load = False - # Receiver's local SourceIdentity, supplied per load_weights() call by - # ModelLoader; the authority for the pre-transfer compatibility gate. + self._transform_protocol_version_for_last_load: Optional[int] = None self._local_source_identity: Optional[SourceIdentity] = None + self._mx_loader = None @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 passed to ModelExpress as - `MODEL_NAME`. The full resolution (with env var and basename - fallbacks) happens inside the transfer and publish paths. - """ - return self._model_name - - @property - def query_timeout_s(self) -> Optional[int]: - return self._query_timeout_s - def is_weights_preloaded(self) -> bool: - """Whether the last :meth:`load_weights` call wired weights directly into the model. - - Reports the result of the most recent `load_weights()` invocation - on this loader instance. `ModelLoader` consults this signal to - decide whether to run the standard weight-mapping pipeline: - - - `True`: MX P2P transfer succeeded; weights already live in - model parameter buffers via direct writes from the upstream - `MxLiveWeightLoader`. The mapping pipeline is skipped for - all parameters covered by P2P. - - `False`: either P2P was never attempted (no MX server URL, - no model reference, library missing) or it failed and we - fell back to disk; weights still need to flow through - `model.load_weights(...)` via the standard mapper. - - Note this is a per-loader-instance flag, not a global one. The - flag is reset to `False` at the start of each `load_weights` - call, so the value is only meaningful immediately after a - successful call. - - Returns: - `True` iff the last `load_weights` populated the model - via P2P; `False` before any call and on any fallback path. - """ return self._p2p_succeeded def is_post_transform_weights_preloaded(self) -> bool: - """Whether the last successful MX preload delivered transformed bytes. - - The source identity bit is included here so callers have one - conservative signal: no identity match, no transform skip. - """ return ( self._p2p_succeeded and self._post_transform_weights_preloaded @@ -319,410 +90,179 @@ def is_post_transform_weights_preloaded(self) -> bool: ) 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. + """Load weights through ModelExpress's shared strategy chain. 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. - `prepare_post_transform_receiver`, when present, is called - after a post-transform source is qualified and before those - direct writes begin. + checkpoint_dir: Hugging Face checkpoint used by native fallback. + mapping: Distributed rank and parallelism mapping. + **kwargs: TRT-LLM-owned load state. ``model``, ``model_config``, + and ``source_identity`` enable the shared strategy path; + ``load_config`` is forwarded to ModelExpress. Qualified + post-transform reception additionally supplies + ``allow_post_transform_weights``, + ``prepare_post_transform_receiver``, and + ``post_transform_protocol_version``. 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. + Native checkpoint weights on fallback, or an empty dictionary + after ModelExpress writes the complete shard into ``model``. + + Raises: + RuntimeError: Qualified reception lacks its structure, + protocol, or current SourceIdentity ABI contract. """ model = kwargs.pop("model", None) - # Popped here so it never leaks into the disk-fallback signature. - self._local_source_identity = kwargs.pop("source_identity", None) + local_source_identity = kwargs.pop("source_identity", None) allow_post_transform_weights = kwargs.pop("allow_post_transform_weights", False) prepare_post_transform_receiver = kwargs.pop("prepare_post_transform_receiver", None) + model_config = kwargs.pop("model_config", None) + load_config = kwargs.pop("load_config", None) + transform_protocol_version = kwargs.pop("post_transform_protocol_version", None) + preserve_mx_session = kwargs.pop("_preserve_mx_session", False) + + missing_mx_state = [ + name + for name, value in ( + ("mx_server_url", self._mx_server_url), + ("model", model), + ("source_identity", local_source_identity), + ("model_config", model_config), + ) + if value is None + ] + if missing_mx_state and preserve_mx_session: + logger.info( + "MX loading skipped for auxiliary native checkpoint load: " + f"missing {', '.join(missing_mx_state)}; preserving the active MX session." + ) + return super().load_weights( + checkpoint_dir, + mapping=mapping, + **kwargs, + ) + + self._local_source_identity = local_source_identity self._p2p_succeeded = False self._post_transform_weights_preloaded = False self._source_identity_compatible_for_last_load = False + self._transform_protocol_version_for_last_load = None + self._cleanup_mx_loader() - if self._mx_server_url is None or model is None: - return self._fallback_to_disk( + if missing_mx_state: + logger.info( + f"MX loading unavailable: missing {', '.join(missing_mx_state)}; " + "falling back to native Hugging Face checkpoint loading." + ) + return super().load_weights( 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)" - ), + mapping=mapping, **kwargs, ) - try: - from modelexpress import ( - trtllm_live_transfer as mx_transfer, # type: ignore[import-not-found] - ) - except ImportError as exc: - raise ImportError( - "ModelExpress checkpoint loading was explicitly requested, " - "but the ModelExpress client could not be imported. Install " - 'the MX dependencies with `pip install "tensorrt-llm[mx]"`, ' - "or select a different " - "`checkpoint_format` to continue without MX." - ) from exc - - # ModelExpress 0.4.1 installs an INFO-level file handler for - # MX_TRANSFER_LOG_DIR, but leaves its logger at Python's WARNING - # default outside vLLM. Enable the records in the worker that performs - # the transfer so the requested per-rank diagnostics are not empty. _enable_mx_transfer_logging() try: - with _MX_TRANSFER_STATE_LOCK: - MxClient = mx_transfer.MxClient - MxLiveWeightLoader = mx_transfer.MxLiveWeightLoader - build_trtllm_identity = mx_transfer._build_trtllm_identity - # Resolve once so discovery and the released ModelExpress - # loader query the same source identity. The lock prevents a - # concurrent MX publish from temporarily changing MODEL_NAME or - # the identity builder while this state is captured. - resolved_name = self._resolve_publish_name(checkpoint_dir) - except AttributeError: - logger.warning( - "modelexpress TRT-LLM live-transfer symbols are missing; " - "cannot use MX P2P weight transfer. Falling back to disk " - "loading." - ) - return self._fallback_to_disk(checkpoint_dir, mapping, **kwargs) - - try: - source_metadata = self._fetch_source_metadata( - checkpoint_dir, - MxClient, - build_trtllm_identity, - model_name=resolved_name, - ) - except Exception: - # Deliberately broad: source discovery is part of the optional MX - # fast path, so an upstream client failure must preserve disk - # loading as the correctness path. - logger.warning( - "MX source metadata fetch failed; falling back to disk " - f"loading.\n{traceback.format_exc()}" - ) - return self._fallback_to_disk( + trtllm_adapter = import_module("modelexpress.engines.trtllm") + except ModuleNotFoundError as exc: + if exc.name == "modelexpress": + logger.warning( + "ModelExpress is not installed; install it with " + '`pip install "tensorrt-llm[mx]"`. Falling back to native ' + "Hugging Face checkpoint loading." + ) + elif exc.name in {"modelexpress.engines", "modelexpress.engines.trtllm"}: + logger.warning( + "The installed ModelExpress package does not provide the TensorRT-LLM " + "adapter; install modelexpress>=0.5.1 or another compatible version. " + "Falling back to native Hugging Face checkpoint loading." + ) + else: + raise + return super().load_weights( checkpoint_dir, - mapping, - reason="MX source metadata probe failed", + mapping=mapping, **kwargs, ) - source_registered = source_metadata is not None - if not source_registered and self._query_timeout_s == 0: - # A zero timeout explicitly disables source polling. Fall back - # before preparing post-transform receiver aliases: that setup - # mutates the module graph and is only safe when P2P will proceed. - return self._fallback_to_disk( - checkpoint_dir, - mapping, - reason="no MX source is registered and source polling is disabled", - **kwargs, - ) - if not source_registered and self._local_source_identity is not None: - # ModelExpress 0.4.1 hashes every SourceIdentity field, including - # extra_parameters. Proceed to MxLiveWeightLoader.load_weights() - # even though this immediate probe found no source: that method - # retries list_sources every five seconds until a source appears or - # query_timeout_s expires. It uses this same patched identity, so - # any source discovered later necessarily carries the expected - # TRT-LLM identity and layout metadata. - source_metadata = _build_mx_source_metadata(self._local_source_identity) - # Pre-transfer compatibility gate: on mismatch, skip the transfer - # before any RDMA work starts and fall back to disk. - self._source_identity_compatible_for_last_load = self._source_metadata_identity_compatible( - source_metadata - ) - if not self._source_identity_compatible_for_last_load: - return self._fallback_to_disk( - checkpoint_dir, - mapping, - reason="source SourceIdentity incompatible with receiver", - **kwargs, + MxModelLoader = getattr(trtllm_adapter, "MxModelLoader", None) + if MxModelLoader is None: + logger.warning( + "The installed ModelExpress TensorRT-LLM adapter is incompatible: " + "MxModelLoader is missing. Install a compatible ModelExpress version. " + "Falling back to native Hugging Face checkpoint loading." ) - - expected_transform_abi_id = ( - self._local_source_identity.transform_abi_id - if self._local_source_identity is not None - else None - ) - layout_status = _metadata_weight_layout_status( - source_metadata, - expected_transform_abi_id=expected_transform_abi_id, - ) - if layout_status is _MxWeightLayoutStatus.UNSUPPORTED: - self._source_identity_compatible_for_last_load = False - return self._fallback_to_disk( + return super().load_weights( checkpoint_dir, - mapping, - reason=_metadata_unsupported_layout_reason( - source_metadata, - expected_transform_abi_id=expected_transform_abi_id, - ), + mapping=mapping, **kwargs, ) - self._post_transform_weights_preloaded = ( - layout_status is _MxWeightLayoutStatus.POST_TRANSFORM_SUPPORTED - ) - if self._post_transform_weights_preloaded and not allow_post_transform_weights: - self._post_transform_weights_preloaded = False - self._source_identity_compatible_for_last_load = False - return self._fallback_to_disk( - checkpoint_dir, - mapping, - reason=( - "source publishes post-transform weights but this model is " - "not qualified for staged MX receiver loading" - ), - **kwargs, + if allow_post_transform_weights and prepare_post_transform_receiver is None: + raise RuntimeError("Qualified MX loading requires receiver structure preparation") + if allow_post_transform_weights and transform_protocol_version is None: + raise RuntimeError("Qualified MX loading requires a transform protocol version") + if allow_post_transform_weights and ( + self._local_source_identity.format_version != SOURCE_IDENTITY_FORMAT_VERSION + or not self._local_source_identity.transform_abi_id + ): + raise RuntimeError( + "Qualified MX loading requires the current TRT-LLM " + "SourceIdentity format and a transform-layout ABI" ) - if self._post_transform_weights_preloaded: - if prepare_post_transform_receiver is None: - self._post_transform_weights_preloaded = False - self._source_identity_compatible_for_last_load = False - return self._fallback_to_disk( - checkpoint_dir, - mapping, - reason=( - "post-transform source requires receiver structure " - "preparation before exact-name P2P transfer" - ), - **kwargs, - ) - # Source-side setup_aliases() may change the first canonical name - # returned for aliased parameters. Mirror that structural state on - # the receiver before upstream MX matches tensors by exact name. - prepare_post_transform_receiver(model) - timeout_override = self._resolve_query_timeout_override( - source_registered=source_registered, - model_name=resolved_name, + mx_loader = MxModelLoader( + model_config=model_config, + load_config=load_config, + checkpoint_loader=self, + checkpoint_dir=checkpoint_dir, + native_loader_kwargs=kwargs, + mapping=mapping, + source_identity=self._local_source_identity, + prepare_post_transform_receiver=( + prepare_post_transform_receiver + if prepare_post_transform_receiver is not None + else lambda _model: None + ), + transform_protocol_version=transform_protocol_version, + p2p_enabled=allow_post_transform_weights, + mx_server_url=self._mx_server_url, ) + self._mx_loader = mx_loader try: - with ( - _MX_TRANSFER_STATE_LOCK, - _temporary_env("MX_SOURCE_QUERY_TIMEOUT", timeout_override), - _temporary_env("MODEL_NAME", resolved_name), - _patched_trtllm_identity_builder(mx_transfer, self._local_source_identity), - ): - 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() + weights = mx_loader.load_model(model) + self._p2p_succeeded = mx_loader.p2p_succeeded + self._transform_protocol_version_for_last_load = mx_loader.transform_protocol_version + if self._p2p_succeeded and weights: + raise RuntimeError("MX P2P loading must not return native checkpoint weights") + post_transform_compatible = ( + self._p2p_succeeded + and transform_protocol_version is not None + and self._transform_protocol_version_for_last_load == transform_protocol_version + # The real ModelExpress TRT-LLM adapter serializes the complete + # authoritative SourceIdentity into its discovery identity. RDMA + # success therefore means the selected source matched format v3, + # the transform-layout ABI, and the remaining TRT identity fields. + and self._local_source_identity.format_version == SOURCE_IDENTITY_FORMAT_VERSION + and bool(self._local_source_identity.transform_abi_id) ) - if self._post_transform_weights_preloaded: - self._post_transform_weights_preloaded = False - self._source_identity_compatible_for_last_load = False - logger.warning( - "MX P2P returned %d fallback weights (%.2f MiB, size mismatch) " - "from a post-transform source at %s. Falling back to a full " - "disk load to avoid mixing transformed P2P tensors with raw " - "fallback tensors before the full post-load transform path.", - len(fallback_weights), - fallback_bytes / (1 << 20), - self._mx_server_url, - ) - return self._fallback_to_disk( - checkpoint_dir, - mapping, - reason="post-transform source returned partial fallback weights", - **kwargs, + if self._p2p_succeeded and not post_transform_compatible: + # Discovery includes the transform protocol in SourceIdentity, so + # this is only a backstop. RDMA already wrote the receiver buffers; + # falling back to disk here could mix transformed and raw weights. + raise RuntimeError( + "MX transferred weights without a compatible TRT-LLM " + "transform protocol and SourceIdentity ABI" ) - # 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 + self._post_transform_weights_preloaded = post_transform_compatible + self._source_identity_compatible_for_last_load = post_transform_compatible + return weights + except Exception: + self._p2p_succeeded = False self._post_transform_weights_preloaded = False - 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, - *, - source_registered: bool, - model_name: str, - ) -> 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 source_registered: - return None - - logger.warning( - "No MX source is currently registered for " - f"{model_name}; " - 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 _source_metadata_identity_compatible(self, metadata: Optional[dict[str, Any]]) -> bool: - source_identity = _source_identity_from_metadata(metadata) - return self._source_identity_compatible_with_source(source_identity) - - def _source_identity_compatible_with_source( - self, source_identity: Optional[SourceIdentity] - ) -> bool: - local_identity = self._local_source_identity - decision = check_weight_sharing_compatibility( - local_identity, - source_identity, - IdentityCheckPolicy.WARN_FALLBACK, - ) - return decision.should_share - - def _fetch_source_metadata( - self, - checkpoint_dir: str, - MxClient: Type[Any], - build_identity: Callable[..., Any], - *, - model_name: Optional[str] = None, - ) -> Optional[dict[str, Any]]: - """Fetch TRT-LLM metadata for the selected MX source, if available.""" - client = None - try: - identity = self._build_mx_identity( - checkpoint_dir, - build_identity, - self._local_source_identity, - model_name=model_name, - ) - client = MxClient(server_url=self._mx_server_url) - for method_name in ("get_source_metadata", "get_metadata", "get_worker_metadata"): - method = getattr(client, method_name, None) - if not callable(method): - continue - try: - metadata = method(identity=identity) - except TypeError: - try: - metadata = method(identity) - except TypeError: - # modelexpress 0.4.1 get_metadata() takes - # mx_source_id/worker_id rather than an identity. Fall - # through to the exact-identity list_sources query. - continue - metadata_dict = _metadata_to_dict(metadata) - if _metadata_has_trtllm_key(metadata_dict): - return metadata_dict - - list_resp = client.list_sources(identity=identity) - instances = _source_instances_from_list_response(list_resp) - metadata_candidates = [] - for instance in instances: - metadata_dict = _source_instance_metadata(instance) - if metadata_dict: - metadata_candidates.append(metadata_dict) - selected_metadata = self._select_source_metadata(metadata_candidates) - if selected_metadata is not None: - return selected_metadata - - # modelexpress 0.4.1 SourceInstanceRef intentionally omits the - # queried SourceIdentity. A non-empty response still proves an - # exact match because list_sources hashes every identity field, - # including extra_parameters. Reconstruct the metadata that was - # embedded in the exact query so the compatibility/layout checks - # remain fail-closed without a second metadata channel. - if instances and self._local_source_identity is not None: - return _build_mx_source_metadata(self._local_source_identity) - return None - finally: - _close_mx_client(client) - - def _build_mx_identity( - self, - checkpoint_dir: str, - build_identity: Callable[..., _MxSourceIdentity], - source_identity: Optional[SourceIdentity], - *, - model_name: Optional[str] = None, - ) -> _MxSourceIdentity: - """Build the MX identity used for discovery and attach TRT-LLM identity.""" - resolved_name = model_name or self._resolve_publish_name(checkpoint_dir) - return _attach_trtllm_metadata_to_mx_identity( - build_identity(model_name=resolved_name), - source_identity, - ) - - def _select_source_metadata( - self, metadata_candidates: list[dict[str, Any]] - ) -> Optional[dict[str, Any]]: - """Select metadata that matches the receiver identity when possible.""" - if not metadata_candidates: - return None - for metadata in metadata_candidates: - if self._source_metadata_matches_local_identity(metadata): - return metadata - return metadata_candidates[0] - - def _source_metadata_matches_local_identity(self, metadata: dict[str, Any]) -> bool: - local_identity = getattr(self, "_local_source_identity", None) - source_identity = _source_identity_from_metadata(metadata) - if local_identity is None or source_identity is None: - return False - return local_identity.matches(source_identity).matched - - 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) + self._source_identity_compatible_for_last_load = False + self._transform_protocol_version_for_last_load = None + self._cleanup_mx_loader() + raise def publish_as_source( self, @@ -731,108 +271,14 @@ def publish_as_source( *, source_identity: Optional[SourceIdentity] = None, ) -> None: - """Publish this instance's weights so other ranks can pull via P2P. - - Called by the integration in `model_loader.py` after - `post_load_weights()` so targets receive the post-transform runtime - layout and, when qualified, can skip their own one-shot transforms. + """Publish through the active MX session. - 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. - source_identity: Source identity built before weight load from - the same lifecycle point on producer and receiver. + ``checkpoint_dir`` and ``source_identity`` are retained only for the + common checkpoint-loader hook signature; the session captured both at + load time. """ - - if self._mx_server_url is None: - return - if source_identity is None: - logger.warning( - "Skipping MX post-transform publish because SourceIdentity is " - "unavailable; receivers cannot safely verify transformed weights." - ) - return - if source_identity.transform_abi_id is None: - logger.warning( - "Skipping MX post-transform publish because SourceIdentity has " - "no qualified transform-layout ABI." - ) - return - - try: - from modelexpress import ( - trtllm_live_transfer as mx_transfer, # type: ignore[import-not-found] - ) - except ImportError: - logger.debug("modelexpress library not installed; skipping MX publish.") - return - try: - publish_model_params = mx_transfer.publish_model_params - except AttributeError: - logger.debug("modelexpress publish_model_params is missing; 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. MX transfer and publish calls in - # this interpreter are serialized while upstream requires process-wide - # state. Tracked as MX-2 in §15 (the env-var dance goes away when - # upstream exports a public identity builder / publish API). - metadata = _build_mx_source_metadata(source_identity) - metadata_kwargs = _publish_metadata_kwargs(publish_model_params, metadata) or {} - identity_builder = getattr(mx_transfer, "_build_trtllm_identity", None) - if not metadata_kwargs and not callable(identity_builder): - logger.warning( - "Skipping MX post-transform publish because " - "publish_model_params does not accept metadata and MX does " - "not expose its TRT-LLM identity builder; receivers cannot " - "safely verify transformed weights." - ) - return - - if threading.active_count() > 1: - logger.warning_once( - "MX publish uses process-wide MODEL_EXPRESS_URL/MODEL_NAME " - "environment variables; concurrent MX transfer and 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", - ) - try: - with _MX_TRANSFER_STATE_LOCK: - resolved_name = self._resolve_publish_name(checkpoint_dir) - # Post-load transforms may enqueue asynchronous writes. Make - # the source buffers globally ready before MX publishes their - # addresses and allows a receiver to issue RDMA reads. - _synchronize_cuda_for_mx_publish() - with ( - _temporary_env("MODEL_EXPRESS_URL", self._mx_server_url), - _temporary_env("MODEL_NAME", resolved_name), - _patched_trtllm_identity_builder(mx_transfer, source_identity), - ): - publish_model_params(model, **metadata_kwargs) - logger.info( - "Published post-transform 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()}" - ) + if self._mx_loader is not None: + self._mx_loader.publish_model(model) def post_load_publish( self, @@ -842,279 +288,23 @@ def post_load_publish( weights_preloaded: bool = False, source_identity: Optional[SourceIdentity] = None, ) -> None: - """Publish locally loaded weights as an MX source when appropriate. - - Args: - model: The loaded model whose parameters should be published for - future MX P2P receivers. - checkpoint_dir: Checkpoint directory used as a fallback model - identity when no explicit MX model name is configured. - weights_preloaded: Whether this worker already received weights - through MX P2P. When true, this worker is an MX receiver and - should not republish the same weights as a source. - source_identity: Producer identity serialized into MX metadata so - receivers can verify layout compatibility before transfer. - - Returns: - None. - """ - if weights_preloaded: - return + """Publish final post-transform weights after any successful load path.""" self.publish_as_source( - model, checkpoint_dir=checkpoint_dir, source_identity=source_identity - ) - - -# --------------------------------------------------------------------------- -# 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. - - 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" - - # 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" - - -def _build_mx_source_metadata(source_identity: Optional[SourceIdentity]) -> dict[str, str]: - metadata = { - _MX_WEIGHT_LAYOUT_METADATA_KEY: _MX_WEIGHT_LAYOUT_POST_TRANSFORM, - _MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY: str(_MX_STAGED_TRANSFORM_PROTOCOL_VERSION), - } - if source_identity is not None: - metadata[_MX_SOURCE_IDENTITY_METADATA_KEY] = _serialize_source_identity(source_identity) - if source_identity.transform_abi_id is not None: - metadata[_MX_TRANSFORM_ABI_ID_METADATA_KEY] = source_identity.transform_abi_id - return metadata - - -def _publish_metadata_kwargs( - publish_model_params: Callable[..., Any], - metadata: dict[str, str], -) -> Optional[dict[str, dict[str, str]]]: - try: - signature = inspect.signature(publish_model_params) - except (TypeError, ValueError): - return None - - parameters = signature.parameters - if "metadata" in parameters: - return {"metadata": metadata} - if "worker_metadata" in parameters: - return {"worker_metadata": metadata} - if any(param.kind == inspect.Parameter.VAR_KEYWORD for param in parameters.values()): - return {"metadata": metadata} - return None - - -def _metadata_to_dict(metadata: Any) -> dict[str, Any]: - if metadata is None or isinstance(metadata, (str, bytes)): - return {} - if isinstance(metadata, dict): - return dict(metadata) - - items = getattr(metadata, "items", None) - if callable(items): - try: - return dict(items()) - except (TypeError, ValueError): - pass - - if type(metadata).__module__.startswith("unittest.mock"): - return {} - - attrs = getattr(metadata, "__dict__", None) - if isinstance(attrs, dict): - return dict(attrs) - return {} - - -def _metadata_get(metadata: Optional[dict[str, Any]], key: str) -> Any: - if not metadata: - return None - return metadata.get(key) - - -def _metadata_has_trtllm_key(metadata: dict[str, Any]) -> bool: - return any( - key in metadata - for key in ( - _MX_SOURCE_IDENTITY_METADATA_KEY, - _MX_WEIGHT_LAYOUT_METADATA_KEY, - _MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY, - _MX_TRANSFORM_ABI_ID_METADATA_KEY, - ) - ) - - -def _source_identity_from_metadata(metadata: Optional[dict[str, Any]]) -> Optional[SourceIdentity]: - value = _metadata_get(metadata, _MX_SOURCE_IDENTITY_METADATA_KEY) - if value is None: - return None - - try: - if isinstance(value, SourceIdentity): - return value - if isinstance(value, bytes): - value = value.decode("utf-8") - if isinstance(value, str): - value = json.loads(value) - to_dict = getattr(value, "to_dict", None) - if callable(to_dict): - value = to_dict() - if not isinstance(value, dict): - raise TypeError(f"expected dict-compatible SourceIdentity, got {type(value)!r}") - return SourceIdentity.from_dict(value) - except (json.JSONDecodeError, KeyError, TypeError, ValueError): - logger.warning( - "MX source metadata contains an invalid SourceIdentity; falling back to disk loading." + model, + checkpoint_dir=checkpoint_dir, + source_identity=source_identity, ) - return None - - -def _metadata_is_post_transform( - metadata: Optional[dict[str, Any]], - *, - expected_transform_abi_id: Optional[str], -) -> bool: - return ( - _metadata_weight_layout_status( - metadata, - expected_transform_abi_id=expected_transform_abi_id, - ) - is _MxWeightLayoutStatus.POST_TRANSFORM_SUPPORTED - ) - - -def _metadata_weight_layout_status( - metadata: Optional[dict[str, Any]], - *, - expected_transform_abi_id: Optional[str], -) -> _MxWeightLayoutStatus: - layout = _metadata_get(metadata, _MX_WEIGHT_LAYOUT_METADATA_KEY) - if layout is None: - return _MxWeightLayoutStatus.PRE_TRANSFORM - - normalized_layout = str(layout).lower() - if normalized_layout == "pre_transform": - return _MxWeightLayoutStatus.PRE_TRANSFORM - if normalized_layout != _MX_WEIGHT_LAYOUT_POST_TRANSFORM: - return _MxWeightLayoutStatus.UNSUPPORTED - - version = _metadata_get(metadata, _MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY) - try: - protocol_version = int(version) - except (TypeError, ValueError): - return _MxWeightLayoutStatus.UNSUPPORTED - if protocol_version != _MX_STAGED_TRANSFORM_PROTOCOL_VERSION: - return _MxWeightLayoutStatus.UNSUPPORTED - source_transform_abi_id = _metadata_get(metadata, _MX_TRANSFORM_ABI_ID_METADATA_KEY) - if not isinstance(source_transform_abi_id, str) or not source_transform_abi_id: - return _MxWeightLayoutStatus.UNSUPPORTED - if expected_transform_abi_id is None or source_transform_abi_id != expected_transform_abi_id: - return _MxWeightLayoutStatus.UNSUPPORTED - return _MxWeightLayoutStatus.POST_TRANSFORM_SUPPORTED + def cleanup(self) -> None: + self._cleanup_mx_loader() + super().cleanup() - -def _metadata_unsupported_layout_reason( - metadata: Optional[dict[str, Any]], - *, - expected_transform_abi_id: Optional[str], -) -> str: - layout = _metadata_get(metadata, _MX_WEIGHT_LAYOUT_METADATA_KEY) - if str(layout).lower() == _MX_WEIGHT_LAYOUT_POST_TRANSFORM: - version = _metadata_get(metadata, _MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY) + def _cleanup_mx_loader(self) -> None: + mx_loader = self._mx_loader + self._mx_loader = None + if mx_loader is None: + return try: - protocol_version = int(version) - except (TypeError, ValueError): - protocol_version = None - if protocol_version != _MX_STAGED_TRANSFORM_PROTOCOL_VERSION: - return ( - "source publishes post-transform weights with unsupported " - f"transform protocol {version!r}" - ) - - source_transform_abi_id = _metadata_get(metadata, _MX_TRANSFORM_ABI_ID_METADATA_KEY) - if not isinstance(source_transform_abi_id, str) or not source_transform_abi_id: - return "source publishes post-transform weights without a transform-layout ABI" - if expected_transform_abi_id is None: - return "receiver has no qualified transform-layout ABI for post-transform weights" - return ( - "source publishes post-transform weights with transform-layout ABI " - f"{source_transform_abi_id!r}; receiver requires {expected_transform_abi_id!r}" - ) - return f"source publishes unsupported MX weight layout {layout!r}" - - -def _source_instances_from_list_response(list_resp: Any) -> list[Any]: - if isinstance(list_resp, dict): - instances = list_resp.get("instances", []) - else: - instances = getattr(list_resp, "instances", []) - return list(instances or []) - - -def _source_instance_metadata(instance: Any) -> dict[str, Any]: - for candidate in ( - instance, - _metadata_attr(instance, "metadata"), - _metadata_attr(instance, "worker_metadata"), - _metadata_attr(instance, "source_metadata"), - ): - metadata = _metadata_to_dict(candidate) - if _metadata_has_trtllm_key(metadata): - return metadata - return {} - - -def _metadata_attr(instance: Any, name: str) -> Any: - if isinstance(instance, dict): - return instance.get(name) - return getattr(instance, name, None) + mx_loader.cleanup() + except Exception as exc: # noqa: BLE001 - cleanup is best effort + logger.warning(f"Failed to clean up ModelExpress loader {mx_loader!r}: {exc!r}") diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 083245faa7a9..37c552fc14fd 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -422,7 +422,6 @@ def __init__( llm_args.checkpoint_loader, llm_args.checkpoint_format, mx_config=llm_args.mx_config, - mx_model_name=llm_args.model, checkpoint_io_policy=llm_args.checkpoint_io_policy, load_format=llm_args.load_format, partial_model_loading=llm_args.is_partial_model_loading, diff --git a/tensorrt_llm/_torch/pyexecutor/model_loader.py b/tensorrt_llm/_torch/pyexecutor/model_loader.py index ccb41552357b..44d41f539324 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_loader.py +++ b/tensorrt_llm/_torch/pyexecutor/model_loader.py @@ -405,7 +405,6 @@ def _construct_checkpoint_loader( checkpoint_format: Optional[str], *, mx_config: Optional[ModelExpressConfig] = None, - mx_model_name: Optional[str] = None, checkpoint_io_policy: str = "native", load_format: LoadFormat | str = LoadFormat.AUTO, partial_model_loading: bool = False, @@ -458,9 +457,6 @@ def _construct_checkpoint_loader( if checkpoint_format == "MX": if mx_config is not None: extra_kwargs["mx_server_url"] = mx_config.server_url - extra_kwargs["query_timeout_s"] = mx_config.server_query_timeout_s - if mx_model_name is not None: - extra_kwargs["model_name"] = mx_model_name checkpoint_loader = BaseCheckpointLoader.get( checkpoint_format=checkpoint_format, @@ -642,6 +638,7 @@ def __init__(self, self.weight_mapper = None self._weight_pool_proxy = None self._gms_backend = None + self._checkpoint_loader: Optional[BaseCheckpointLoader] = None # Mostly weight loading and processing time metrics, updated when load() is called. self._metrics: dict[str, float] = {} @@ -796,6 +793,7 @@ def load( The loaded and initialized PyTorch model. """ self._metrics = {} + self._checkpoint_loader = checkpoint_loader config = self._load_and_validate_config(checkpoint_dir, checkpoint_loader) # Some model constructors normalize or rewrite config fields. Capture @@ -949,6 +947,10 @@ def init_meta_tensor(t: torch.Tensor): "source_identity": self._source_identity, } if checkpoint_loader.checkpoint_format == "MX": + load_weights_kwargs["model_config"] = config + load_weights_kwargs["load_config"] = self.llm_args + load_weights_kwargs[ + "post_transform_protocol_version"] = self._MX_STAGED_RECEIVER_TRANSFORM_PROTOCOL_VERSION # If a separate draft model still needs a raw disk load, # do not accept post-transform bytes for only the target # model. Enable this only after target and draft subgraphs @@ -1052,6 +1054,11 @@ def init_meta_tensor_in_pool(t: torch.Tensor): "source_identity": self._source_identity, } if checkpoint_loader.checkpoint_format == "MX": + load_weights_kwargs["model_config"] = config + load_weights_kwargs[ + "load_config"] = self.llm_args + load_weights_kwargs[ + "post_transform_protocol_version"] = self._MX_STAGED_RECEIVER_TRANSFORM_PROTOCOL_VERSION load_weights_kwargs[ "allow_post_transform_weights"] = post_transform_qualification.qualified if post_transform_qualification.qualified: @@ -1162,12 +1169,9 @@ def init_meta_tensor_in_pool(t: torch.Tensor): # RO path: weights are coming from a GMS donor that # has already committed the post-post_load layout, so # the receiver flags `weights_preloaded=True` on the - # checkpoint-loader hooks. `MXCheckpointLoader.post_load_publish` - # (and any other format-aware loader) honors this flag - # to early-return and not re-publish, while still - # letting `post_load_apply` perform any - # receiver-side per-format work (e.g., marking - # presharded modules). + # checkpoint-loader hooks. This prevents duplicate + # transforms while leaving any receiver-side publication + # policy to the format-specific hook. # # Hook order: # 1. `post_load_apply`: format-specific apply @@ -1183,8 +1187,8 @@ def init_meta_tensor_in_pool(t: torch.Tensor): # 5. Per-module `cache_derived_state`: recompute # Python-side state from real, materialized # tensors without re-running one-shot transforms. - # 6. `post_load_publish`: any receiver-side - # publish (no-op via the receiver guard). + # 6. `post_load_publish`: apply the format-specific + # receiver publication policy. with timing_metric( ModelLoaderMetricNames. POST_LOAD_PROCESSING_SECONDS.value, @@ -1397,11 +1401,19 @@ def _load_separate_draft_weights( One-model MTP with separate heads reuses the target architecture mapper because MTP modules are already attached under the target model. """ + draft_load_kwargs = {} + if checkpoint_loader.checkpoint_format == "MX": + # This loader instance still owns the target model's MX session. + # Preserve it while the native loader reads the auxiliary draft + # checkpoint so the finalized target can publish through it. + draft_load_kwargs["_preserve_mx_session"] = True with timing_metric( ModelLoaderMetricNames.DRAFT_CHECKPOINT_PREPARATION_SECONDS. value, self._metrics): draft_weights = checkpoint_loader.load_weights( - self.spec_config.speculative_model, mapping=self.mapping) + self.spec_config.speculative_model, + mapping=self.mapping, + **draft_load_kwargs) if model.draft_config is not None: draft_model_arch = model.draft_config.pretrained_config.architectures[ @@ -1647,22 +1659,31 @@ def abort_update_weights(self) -> None: def cleanup(self) -> None: """Release backend resources acquired during :meth:`load`. - Currently the only backend held by `ModelLoader` is the - optional GMS client, established by the `LoadFormat.GMS` - branch. Releasing it disconnects from the GMS daemon and evicts - the per-tag client registry entry; weights remain alive - on-device for any other process holding an RO lock on the same - `tag`. - - Idempotent: a second call after a successful cleanup is a no-op - because the backend handle is dropped. Best-effort: any failure - in the underlying `GMSBackend.cleanup()` is swallowed there - and logged, so this method never raises — safe to call from - :meth:`PyTorchModelEngine.cleanup` and `__del__` paths. + This releases the optional GMS client and the active checkpoint + loader. Releasing GMS disconnects from the daemon and evicts the + per-tag client registry entry; weights remain alive on-device for + any other process holding an RO lock on the same `tag`. + + A second call is a no-op because both handles are dropped. Cleanup is + best effort so shutdown and destructor paths do not propagate backend + failures. """ if self._gms_backend is not None: - self._gms_backend.cleanup() + gms_backend = self._gms_backend self._gms_backend = None + try: + gms_backend.cleanup() + except Exception as exc: # noqa: BLE001 - shutdown must remain best effort + logger.warning( + f"Failed to clean up GMS backend {gms_backend!r}: {exc!r}") + if self._checkpoint_loader is not None: + try: + self._checkpoint_loader.cleanup() + except Exception as exc: # noqa: BLE001 - shutdown must remain best effort + logger.warning(f"Failed to clean up checkpoint loader " + f"{self._checkpoint_loader!r}: {exc!r}") + finally: + self._checkpoint_loader = None def _load_and_validate_config( self, checkpoint_dir: str, diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 7674e49b83d7..5023f5cf8fa5 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -259,7 +259,6 @@ def _load_config_and_create_checkpoint_loader( llm_args.checkpoint_loader, llm_args.checkpoint_format, mx_config=llm_args.mx_config, - mx_model_name=llm_args.model, checkpoint_io_policy=llm_args.checkpoint_io_policy, load_format=llm_args.load_format, partial_model_loading=llm_args.is_partial_model_loading, diff --git a/tensorrt_llm/executor/base_worker.py b/tensorrt_llm/executor/base_worker.py index 5d9274da8ea3..4177e41ba7f6 100644 --- a/tensorrt_llm/executor/base_worker.py +++ b/tensorrt_llm/executor/base_worker.py @@ -202,7 +202,6 @@ def _create_py_executor(): self.llm_args.checkpoint_loader, self.llm_args.checkpoint_format, mx_config=self.llm_args.mx_config, - mx_model_name=self.llm_args.model, checkpoint_io_policy=self.llm_args.checkpoint_io_policy, load_format=self.llm_args.load_format, partial_model_loading=partial_model_loading, diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 343cef0fbd63..59a8646314db 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -5212,11 +5212,11 @@ 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.", - status="prototype") + description="Deprecated and ignored. ModelExpress performs one source " + "discovery query and falls back to native checkpoint loading when no " + "compatible source is immediately available.", + status="deprecated", + deprecated=True) preshard_strategy: str = Field( default="per_module", diff --git a/tests/integration/defs/model_express/test_model_express.py b/tests/integration/defs/model_express/test_model_express.py index 158c88a348bf..202247fe811d 100644 --- a/tests/integration/defs/model_express/test_model_express.py +++ b/tests/integration/defs/model_express/test_model_express.py @@ -49,10 +49,9 @@ "sourceidentity mismatch", "invalid sourceidentity", ) -_MATCHED_PARAMS_PATTERN = re.compile(r"Matched\s+(\d+)/(\d+)\s+params", re.IGNORECASE) -_RANK_LOG_PATTERN = re.compile(r"rank(\d+)\.log", re.IGNORECASE) -_TRANSFERRED_PARAMS_PATTERN = re.compile( - r"Rank\s+(\d+):\s+transferred\s+(\d+)\s+params", +_RDMA_TRANSFER_PATTERN = re.compile( + r"\[Worker\s+(\d+)\].*?RDMA transfer complete:\s+" + r"(\d+)\s+tensors,\s+([0-9.]+)\s+GB", re.IGNORECASE, ) _DONOR_PROCESS_FAILURE_MARKERS = ( @@ -61,7 +60,7 @@ b"process returned a non-zero exit code", ) _DONOR_PROCESS_FAILURE_OVERLAP = max(len(marker) for marker in _DONOR_PROCESS_FAILURE_MARKERS) - 1 -_MINIMUM_MODELEXPRESS_VERSION = Version("0.4.1") +_MINIMUM_MODELEXPRESS_VERSION = Version("0.5.1") _MODELEXPRESS_VERSION_PREFIX = "MODELEXPRESS_VERSION=" _MX_PREFLIGHT_SCRIPT = """ import importlib.metadata as metadata @@ -75,9 +74,7 @@ def module_exists(name): return False -assert module_exists("modelexpress.trtllm_live_transfer") or module_exists( - "modelexpress.engines.trtllm" -), "no TRT-LLM adapter is available" +assert module_exists("modelexpress.engines.trtllm"), "no TRT-LLM adapter is available" from modelexpress.nixl_transfer import is_nixl_available assert is_nixl_available(), "NIXL is unavailable" @@ -456,70 +453,36 @@ def _load_tokens(output_path: Path) -> list[list[int]]: return token_ids -def _transfer_logs_by_rank(transfer_log_dir: Path, receiver_log: str) -> dict[int, str]: - log_files = tuple(path for path in transfer_log_dir.rglob("*") if path.is_file()) - logs = tuple(path for path in log_files if path.stat().st_size > 0) - if not logs: - entries = ", ".join( - f"{path.relative_to(transfer_log_dir)} ({path.stat().st_size} bytes)" - for path in log_files - ) - pytest.fail( - f"ModelExpress created no non-empty receiver transfer logs in " - f"{transfer_log_dir}; entries: {entries or ''}\n" - f"Receiver log:\n{receiver_log}" - ) - - logs_by_rank = {} - for path in sorted(logs): - match = _RANK_LOG_PATTERN.fullmatch(path.name) - if match is None: - pytest.fail(f"Unexpected ModelExpress receiver transfer log: {path}") - rank = int(match.group(1)) - if rank in logs_by_rank: - pytest.fail(f"ModelExpress created multiple receiver transfer logs for rank {rank}") - logs_by_rank[rank] = path.read_text(encoding="utf-8", errors="replace") - return logs_by_rank - - def _assert_transfer_evidence( case: MxE2ECase, receiver_log_path: Path, - receiver_transfer_log_dir: Path, ) -> None: receiver_log = receiver_log_path.read_text(encoding="utf-8", errors="replace") - transfer_logs = _transfer_logs_by_rank(receiver_transfer_log_dir, receiver_log) expected_ranks = set(range(case.tp_size)) - assert set(transfer_logs) == expected_ranks, ( - f"Expected receiver transfer logs for ranks {expected_ranks}, got {set(transfer_logs)}" + transfers_by_rank: dict[int, list[tuple[int, float]]] = {} + for rank, tensor_count, size_gb in _RDMA_TRANSFER_PATTERN.findall(receiver_log): + transfers_by_rank.setdefault(int(rank), []).append((int(tensor_count), float(size_gb))) + + assert set(transfers_by_rank) == expected_ranks, ( + f"Expected RDMA transfer completion for ranks {expected_ranks}, " + f"got {set(transfers_by_rank)}\nReceiver log:\n{receiver_log}" ) - all_receiver_logs = receiver_log + "\n" + "\n".join(transfer_logs.values()) - all_receiver_logs_lower = all_receiver_logs.lower() + receiver_log_lower = receiver_log.lower() for marker in _RECEIVER_FAILURE_MARKERS: - assert marker not in all_receiver_logs_lower, ( - f"MX receiver logs contain failure marker {marker!r}" + assert marker not in receiver_log_lower, ( + f"MX receiver log contains failure marker {marker!r}" ) for rank in sorted(expected_ranks): - rank_log = transfer_logs[rank] - matched_params = _MATCHED_PARAMS_PATTERN.findall(rank_log) - assert len(matched_params) == 1, ( - f"Expected one matched-parameter summary for rank {rank}, got {matched_params}" + transfers = transfers_by_rank[rank] + assert len(transfers) == 1, ( + f"Expected one RDMA transfer completion for rank {rank}, got {transfers}" ) - matched, total = (int(value) for value in matched_params[0]) - assert matched == total > 0, ( - f"MX receiver rank {rank} reported incomplete parameter match {matched}/{total}" - ) - - transferred_params = _TRANSFERRED_PARAMS_PATTERN.findall(rank_log) - assert len(transferred_params) == 1, ( - f"Expected one transfer summary for rank {rank}, got {transferred_params}" - ) - transferred_rank, transferred_count = (int(value) for value in transferred_params[0]) - assert transferred_rank == rank and transferred_count == matched, ( - f"MX receiver rank {rank} matched {matched} params but reported transfer summary " - f"{transferred_params[0]}" + tensor_count, size_gb = transfers[0] + assert tensor_count > 0 and size_gb > 0, ( + f"MX receiver rank {rank} reported an empty transfer: " + f"{tensor_count} tensors, {size_gb} GB" ) @@ -530,6 +493,9 @@ def test_mx_donor_receiver(case: MxE2ECase, tmp_path: Path) -> None: mx_url, gpu_ids = _require_mx_environment(required_gpus) model_path = _resolve_model_path(case) timeout_s = int(os.environ.get("TRTLLM_MX_E2E_TIMEOUT_S", "1200")) + # NIXL binds base_port + device_id. The donor and receiver share one CI + # network namespace, so give their TP ranks adjacent, non-overlapping ranges. + metadata_port = int(os.environ.get("MX_METADATA_PORT", "5555")) donor_snapshot = _build_canonical_snapshot(case, model_path, tmp_path) receiver_snapshot = _build_metadata_only_snapshot(donor_snapshot, tmp_path) @@ -563,6 +529,7 @@ def test_mx_donor_receiver(case: MxE2ECase, tmp_path: Path) -> None: ) donor_environment = _worker_environment(donor_gpu_ids, tmp_path / "donor-transfer-logs") + donor_environment["MX_METADATA_PORT"] = str(metadata_port) donor_returncode = None with ( donor_log.open("w", encoding="utf-8") as donor_log_file, @@ -585,6 +552,10 @@ def test_mx_donor_receiver(case: MxE2ECase, tmp_path: Path) -> None: ): try: _wait_for_donor(donor_process, donor_ready, donor_log, timeout_s) + receiver_environment = _worker_environment( + receiver_gpu_ids, tmp_path / "receiver-transfer-logs" + ) + receiver_environment["MX_METADATA_PORT"] = str(metadata_port + case.tp_size) _run_worker( _worker_command( role="receiver", @@ -593,7 +564,7 @@ def test_mx_donor_receiver(case: MxE2ECase, tmp_path: Path) -> None: output_path=receiver_output, mx_url=mx_url, ), - _worker_environment(receiver_gpu_ids, tmp_path / "receiver-transfer-logs"), + receiver_environment, receiver_log, timeout_s, ) @@ -611,5 +582,4 @@ def test_mx_donor_receiver(case: MxE2ECase, tmp_path: Path) -> None: _assert_transfer_evidence( case, receiver_log, - tmp_path / "receiver-transfer-logs", ) diff --git a/tests/integration/test_lists/test-db/l0_cpu.yml b/tests/integration/test_lists/test-db/l0_cpu.yml index 1fa059baca8d..817bf2f4b9d4 100644 --- a/tests/integration/test_lists/test-db/l0_cpu.yml +++ b/tests/integration/test_lists/test-db/l0_cpu.yml @@ -101,6 +101,7 @@ l0_cpu: - unittest/llmapi/test_whisper_suppress_tokens_processor.py - unittest/llmapi/test_kv_cache_dtype_override.py - unittest/llmapi/test_llm_args.py + - unittest/llmapi/test_mx_args.py - unittest/llmapi/test_llm_quant.py - unittest/llmapi/test_llm_telemetry.py - unittest/llmapi/test_llm_utils.py diff --git a/tests/integration/test_lists/test-db/l0_model_express.yml b/tests/integration/test_lists/test-db/l0_model_express.yml index 6a9c3cb5df64..5921d1d3f674 100644 --- a/tests/integration/test_lists/test-db/l0_model_express.yml +++ b/tests/integration/test_lists/test-db/l0_model_express.yml @@ -22,6 +22,7 @@ l0_model_express: - model_express/test_model_express.py::test_mx_donor_receiver[qwen2-bf16-tp1] - model_express/test_model_express.py::test_mx_donor_receiver[qwen3-bf16-tp1] - model_express/test_model_express.py::test_mx_donor_receiver[mistral-bf16-tp1] + - unittest/_torch/weight_sharing/test_mx_source_identity_gate.py - condition: ranges: system_gpu_count: diff --git a/tests/unittest/_torch/executor/test_model_loader_mx.py b/tests/unittest/_torch/executor/test_model_loader_mx.py index fbac9cd1cc3c..7879e601d992 100644 --- a/tests/unittest/_torch/executor/test_model_loader_mx.py +++ b/tests/unittest/_torch/executor/test_model_loader_mx.py @@ -218,6 +218,30 @@ def _llama_alias_state(model): } +def _mx_canonical_parameter_catalog(model: nn.Module): + """Mirror MX's canonical TRT-LLM parameter view without importing MX.""" + catalog = {} + seen_storages = set() + runtime_alias_components = {"next_attn", "next_layer_layernorm"} + for name, parameter in model.named_parameters(remove_duplicate=False): + storage = ( + parameter.device.type, + parameter.device.index, + parameter.data_ptr(), + ) + if runtime_alias_components.intersection(name.split(".")): + continue + if storage in seen_storages: + continue + seen_storages.add(storage) + catalog[name] = ( + parameter.data_ptr(), + tuple(parameter.shape), + parameter.dtype, + ) + return catalog + + def _llama_input_embeddings(model: nn.Module) -> torch.Tensor: input_ids = torch.tensor( [0, 1, 2], @@ -583,13 +607,10 @@ def test_construct_checkpoint_loader_passes_mx_config(): None, "MX", mx_config=mx_config, - mx_model_name="Qwen/Qwen3-8B", ) assert isinstance(checkpoint_loader, MXCheckpointLoader) assert checkpoint_loader.mx_server_url == "http://mx:8001" - assert checkpoint_loader.query_timeout_s == 17 - assert checkpoint_loader.model_name == "Qwen/Qwen3-8B" def _format_documented_values( @@ -698,6 +719,15 @@ def test_mx_success_initializes_mapper_skips_weight_mapping_and_reload_works( assert kwargs["mapping"] is loader.mapping assert kwargs["model"] is model assert kwargs["source_identity"] is loader._source_identity + assert ( + kwargs["model_config"] + is model_loader_mod.AutoModelForCausalLM.from_config.call_args.args[0] + ) + assert kwargs["load_config"] is loader.llm_args + assert ( + kwargs["post_transform_protocol_version"] + == ModelLoader._MX_STAGED_RECEIVER_TRANSFORM_PROTOCOL_VERSION + ) assert kwargs["allow_post_transform_weights"] is True assert loader._source_identity.transform_abi_id == LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1 assert loader._call_load_weights.call_count == 0 @@ -721,6 +751,73 @@ def test_mx_success_initializes_mapper_skips_weight_mapping_and_reload_works( assert events == ["post_load_weights", "load_weights"] +def test_cleanup_releases_active_checkpoint_loader(monkeypatch): + loader = _make_loader(monkeypatch, events=[]) + checkpoint_loader = MagicMock(name="checkpoint_loader") + checkpoint_loader.checkpoint_format = "MX" + checkpoint_loader.is_weights_preloaded.return_value = False + checkpoint_loader.load_weights.return_value = {"weight": MagicMock()} + + loader.load("/ckpt", checkpoint_loader) + loader.cleanup() + + checkpoint_loader.cleanup.assert_called_once_with() + assert loader._checkpoint_loader is None + + +def test_cleanup_swallows_checkpoint_loader_failure(monkeypatch): + loader = _make_loader(monkeypatch, events=[]) + checkpoint_loader = MagicMock(name="checkpoint_loader") + cleanup_error = RuntimeError("cleanup failed") + checkpoint_loader.cleanup.side_effect = cleanup_error + loader._checkpoint_loader = checkpoint_loader + warning = MagicMock() + monkeypatch.setattr(model_loader_mod.logger, "warning", warning) + + loader.cleanup() + + checkpoint_loader.cleanup.assert_called_once_with() + assert loader._checkpoint_loader is None + warning.assert_called_once_with( + f"Failed to clean up checkpoint loader {checkpoint_loader!r}: {cleanup_error!r}" + ) + + +def test_cleanup_continues_after_gms_failure(monkeypatch): + loader = _make_loader(monkeypatch, events=[]) + gms_backend = MagicMock(name="gms_backend") + cleanup_error = RuntimeError("gms cleanup failed") + gms_backend.cleanup.side_effect = cleanup_error + checkpoint_loader = MagicMock(name="checkpoint_loader") + loader._gms_backend = gms_backend + loader._checkpoint_loader = checkpoint_loader + warning = MagicMock() + monkeypatch.setattr(model_loader_mod.logger, "warning", warning) + + loader.cleanup() + + gms_backend.cleanup.assert_called_once_with() + checkpoint_loader.cleanup.assert_called_once_with() + assert loader._gms_backend is None + assert loader._checkpoint_loader is None + warning.assert_called_once_with( + f"Failed to clean up GMS backend {gms_backend!r}: {cleanup_error!r}" + ) + + +def test_config_failure_retains_checkpoint_loader_for_cleanup(monkeypatch): + loader = _make_loader(monkeypatch, events=[]) + checkpoint_loader = MagicMock(name="checkpoint_loader") + loader._load_and_validate_config.side_effect = RuntimeError("config failed") + + with pytest.raises(RuntimeError, match="config failed"): + loader.load("/ckpt", checkpoint_loader) + + assert loader._checkpoint_loader is checkpoint_loader + loader.cleanup() + checkpoint_loader.cleanup.assert_called_once_with() + + @pytest.mark.cpu_only def test_reload_partial_loading_preserves_weights_transformed_flags(monkeypatch): events = [] @@ -823,6 +920,15 @@ def test_mx_post_transform_receiver_uses_staged_path_when_qualified( loader._call_load_weights.assert_not_called() _args, kwargs = checkpoint_loader.load_weights.call_args assert kwargs["allow_post_transform_weights"] is True + assert ( + kwargs["model_config"] + is model_loader_mod.AutoModelForCausalLM.from_config.call_args.args[0] + ) + assert kwargs["load_config"] is loader.llm_args + assert ( + kwargs["post_transform_protocol_version"] + == ModelLoader._MX_STAGED_RECEIVER_TRANSFORM_PROTOCOL_VERSION + ) assert callable(kwargs["prepare_post_transform_receiver"]) checkpoint_loader.post_load_publish.assert_called_once_with( model, @@ -1537,6 +1643,25 @@ def test_bf16_dense_profiles_ignore_moe_only_runtime_dimensions( assert decision.qualified +def test_staged_llama_finalization_preserves_mx_tensor_catalog( + monkeypatch: pytest.MonkeyPatch, +) -> None: + model = _tiny_llama_model(monkeypatch) + + # MX discovers and registers tensors after receiver alias preparation. + ModelLoader._setup_aliases(model) + registered_catalog = _mx_canonical_parameter_catalog(model) + + # This is the receiver-side finalization sequence after RDMA completes. + ModelLoader._setup_aliases(model) + ModelLoader._mark_weights_transformed(model) + ModelLoader._walk_cache_state(model) + published_catalog = _mx_canonical_parameter_catalog(model) + + assert registered_catalog + assert published_catalog == registered_catalog + + @pytest.mark.cpu_only def test_separate_draft_model_is_not_qualified_by_target_only_profile( monkeypatch: pytest.MonkeyPatch, @@ -1558,6 +1683,45 @@ def test_separate_draft_model_is_not_qualified_by_target_only_profile( assert decision.unsupported_features == frozenset({PostTransformFeature.SEPARATE_DRAFT_MODEL}) +@pytest.mark.cpu_only +def test_separate_draft_load_preserves_mx_session( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mapping = MagicMock() + spec_config = SimpleNamespace(speculative_model="draft-checkpoint") + loader = ModelLoader( + llm_args=MagicMock(), + mapping=mapping, + spec_config=spec_config, + sparse_attention_config=None, + max_num_tokens=1, + max_seq_len=1, + ) + checkpoint_loader = MagicMock() + checkpoint_loader.checkpoint_format = "MX" + checkpoint_loader.load_weights.return_value = {"draft": object()} + draft_mapper = MagicMock() + monkeypatch.setattr( + model_loader_mod.AutoCheckpointMapper, + "get", + MagicMock(return_value=draft_mapper), + ) + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + model = SimpleNamespace( + draft_config=_make_draft_model_config(), + draft_model=MagicMock(), + load_draft_weights=lambda _weights: None, + ) + + loader._load_separate_draft_weights(model, checkpoint_loader) + + checkpoint_loader.load_weights.assert_called_once_with( + "draft-checkpoint", + mapping=mapping, + _preserve_mx_session=True, + ) + + @pytest.mark.cpu_only def test_one_engine_speculative_mode_is_not_qualified_by_target_only_profile( monkeypatch: pytest.MonkeyPatch, 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 636f80905864..e0abd75c0a72 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 @@ -6,19 +6,18 @@ # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 -"""Unit tests for MXCheckpointLoader with checkpoint_format='MX'. - -These tests intentionally do not exercise the upstream modelexpress library. -The import-failure path blocks modelexpress symbols from sys.modules so the -assertion is about our dependency handling, not the upstream API. -""" +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Unit tests for the native TRT-LLM entrypoint into MX strategies.""" -import json import logging -import os import sys -from contextlib import ExitStack -from types import SimpleNamespace +from dataclasses import replace +from types import ModuleType from unittest.mock import MagicMock, patch import pytest @@ -30,20 +29,8 @@ Qwen3NextHfWeightMapper, ) from tensorrt_llm._torch.models.checkpoints.hf.weight_mapper import HfWeightMapper -from tensorrt_llm._torch.models.checkpoints.mx import checkpoint_loader as mx_checkpoint_loader -from tensorrt_llm._torch.models.checkpoints.mx.checkpoint_loader import ( - _MX_SOURCE_IDENTITY_METADATA_KEY, - _MX_STAGED_TRANSFORM_PROTOCOL_VERSION, - _MX_TRANSFORM_ABI_ID_METADATA_KEY, - _MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY, - _MX_WEIGHT_LAYOUT_METADATA_KEY, - _MX_WEIGHT_LAYOUT_POST_TRANSFORM, - MXCheckpointLoader, - _build_mx_source_metadata, - _normalize_model_identity, - _resolve_mx_model_name, - _serialize_source_identity, -) +from tensorrt_llm._torch.models.checkpoints.mx import checkpoint_loader as checkpoint_loader_mod +from tensorrt_llm._torch.models.checkpoints.mx.checkpoint_loader import MXCheckpointLoader from tensorrt_llm._torch.weight_sharing import ( ARTIFACT_IDENTITY_FORMAT_VERSION, LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1, @@ -52,15 +39,10 @@ SourceIdentity, ) -_MISSING = object() +pytestmark = pytest.mark.cpu_only -def _identity( - rank: int = 0, - suffix: str = "same", - *, - transform_abi_id: str | None = LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1, -) -> SourceIdentity: +def _source_identity() -> SourceIdentity: return SourceIdentity( format_version=SOURCE_IDENTITY_FORMAT_VERSION, artifact_identity=ArtifactIdentity( @@ -68,1327 +50,578 @@ def _identity( scheme="checkpoint_manifest_sha256", digest="0" * 64, ), - model_fingerprint=f"model-{suffix}", - quant_fingerprint=f"quant-{suffix}", - backend_fingerprint=f"backend-{suffix}", - parallel_fingerprint=f"parallel-{suffix}", - rank=rank, - shard_fingerprint=f"shard-{rank}-{suffix}", - transform_abi_id=transform_abi_id, - model_name="TinyLlama/TinyLlama-1.1B-Chat-v1.0", + model_fingerprint="model", + quant_fingerprint="quant", + backend_fingerprint="backend", + parallel_fingerprint="parallel", + rank=0, + shard_fingerprint="shard", + transform_abi_id=LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1, + model_name="meta-llama/Llama-3.1-8B-Instruct", ) -def _source_identity(rank=0, suffix="same"): - return _identity(rank=rank, suffix=suffix) - - -def _source_instance( - identity: SourceIdentity | None = None, *, post_transform: bool = True, rank=0 -): - if identity is not None: - metadata = _build_mx_source_metadata(identity) - if not post_transform: - metadata[_MX_WEIGHT_LAYOUT_METADATA_KEY] = "pre_transform" - return SimpleNamespace(metadata=metadata, worker_rank=identity.rank) - return SimpleNamespace( - mx_source_id=f"source-{rank}", - worker_id=f"worker-{rank}", - worker_rank=rank, +def _loader(**kwargs): + weight_loader = MagicMock() + weight_loader.load_weights.return_value = {"disk": object()} + config_loader = MagicMock() + loader = MXCheckpointLoader( + weight_loader=weight_loader, + config_loader=config_loader, + **kwargs, ) + return loader, weight_loader, config_loader + + +def _load_kwargs(**overrides): + values = { + "mapping": MagicMock(), + "model": MagicMock(), + "source_identity": _source_identity(), + "model_config": MagicMock(), + "load_config": MagicMock(), + "allow_post_transform_weights": True, + "prepare_post_transform_receiver": MagicMock(), + "post_transform_protocol_version": 1, + } + values.update(overrides) + return values + + +def _install_fake_mx( + monkeypatch, + *, + p2p_succeeded, + value, + transform_protocol_version=1, + load_error=None, +): + instances = [] + + class MxModelLoader: + def __init__(self, **kwargs): + self.kwargs = kwargs + self.p2p_succeeded = p2p_succeeded + self.transform_protocol_version = transform_protocol_version if p2p_succeeded else None + self.publish_model = MagicMock() + self.cleanup = MagicMock() + instances.append(self) + + def load_model(self, model): + self.model = model + if load_error is not None: + raise load_error + return value + + module = ModuleType("modelexpress.engines.trtllm") + module.MxModelLoader = MxModelLoader + monkeypatch.setitem(sys.modules, module.__name__, module) + return instances + + +def test_construction_preserves_checkpoint_loader_contract(): + loader, _, _ = _loader(mx_server_url="mx:8001") + + assert isinstance(loader, HfCheckpointLoader) + assert isinstance(loader, BaseCheckpointLoader) + assert loader.checkpoint_format == "MX" + assert loader._checkpoint_format == "MX" + assert loader.mx_server_url == "mx:8001" + assert not loader.is_weights_preloaded() + + +@pytest.mark.parametrize( + ("effective_level", "expected_level"), + ((logging.WARNING, logging.INFO), (logging.DEBUG, None)), +) +def test_transfer_log_dir_enables_info_records(monkeypatch, effective_level, expected_level): + monkeypatch.setenv("MX_TRANSFER_LOG_DIR", "/tmp/mx-transfer-logs") + mx_logger = MagicMock() + mx_logger.getEffectiveLevel.return_value = effective_level + with patch.object(checkpoint_loader_mod.logging, "getLogger", return_value=mx_logger): + checkpoint_loader_mod._enable_mx_transfer_logging() -pytestmark = pytest.mark.cpu_only + if expected_level is None: + mx_logger.setLevel.assert_not_called() + else: + mx_logger.setLevel.assert_called_once_with(expected_level) -# --------------------------------------------------------------------------- -# Construction & static properties -# --------------------------------------------------------------------------- - - -class TestConstruction: - def test_no_args_constructs(self): - loader = MXCheckpointLoader() - assert loader.mx_server_url is None - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() 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.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - - def test_query_timeout_stored(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001", query_timeout_s=900) - assert loader.query_timeout_s == 900 - - 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) - - def test_checkpoint_format_property(self): - loader = MXCheckpointLoader() - assert loader.checkpoint_format == "MX" - - def test_checkpoint_format_backing_attr(self): - # Some call sites read self._checkpoint_format directly. Keep the - # backing attribute aligned with the property override. - loader = MXCheckpointLoader() - assert loader._checkpoint_format == "MX" - - def test_is_weights_preloaded_initial(self): - loader = MXCheckpointLoader() - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - - def test_post_transform_signal_requires_p2p_and_identity_match(self): - loader = MXCheckpointLoader() - loader._p2p_succeeded = True - loader._post_transform_weights_preloaded = True - loader._source_identity_compatible_for_last_load = False - assert loader.is_post_transform_weights_preloaded() is False - - loader._source_identity_compatible_for_last_load = True - assert loader.is_post_transform_weights_preloaded() is True - - @pytest.mark.parametrize( - ("effective_level", "expected_level"), - ((logging.WARNING, logging.INFO), (logging.DEBUG, None)), +def test_registered_under_mx_and_mapper_fallback_is_preserved(): + loader = BaseCheckpointLoader.get( + checkpoint_format="MX", + weight_loader=None, + weight_mapper=None, + config_loader=None, + mx_server_url="mx:8001", ) - def test_transfer_log_dir_enables_info_records( - self, monkeypatch, effective_level, expected_level - ): - monkeypatch.setenv("MX_TRANSFER_LOG_DIR", "/tmp/mx-transfer-logs") - mx_logger = MagicMock() - mx_logger.getEffectiveLevel.return_value = effective_level - - with patch.object(mx_checkpoint_loader.logging, "getLogger", return_value=mx_logger): - mx_checkpoint_loader._enable_mx_transfer_logging() - - if expected_level is None: - mx_logger.setLevel.assert_not_called() - else: - mx_logger.setLevel.assert_called_once_with(expected_level) - - -# --------------------------------------------------------------------------- -# Registry -# --------------------------------------------------------------------------- - - -class TestRegistry: - def test_registered_under_mx(self): - # BaseCheckpointLoader.get resolves MX through the loader registry. - # Use the same constructor shape as _construct_checkpoint_loader. - 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 with no upstream library involved. -# --------------------------------------------------------------------------- + assert isinstance(loader, MXCheckpointLoader) + assert isinstance( + AutoCheckpointMapper.get("MX", "Qwen3NextForCausalLM"), + Qwen3NextHfWeightMapper, + ) + assert isinstance( + AutoCheckpointMapper.get("MX", "UnknownArchitecture"), + HfWeightMapper, + ) -class TestLoadWeightsFallback: - """Disk-fallback paths that should not touch the upstream MX library. - - All fallback triggers share the same observable contract: - is_weights_preloaded() stays False, HfCheckpointLoader.load_weights is - invoked exactly once, and its return value is propagated unchanged. - """ - # Trigger setup builders. - @staticmethod - def _no_url(stack): # noqa: ARG004 - stack unused for this trigger. - return MXCheckpointLoader(), {"model": MagicMock()} +@pytest.mark.parametrize( + ("loader_kwargs", "load_overrides", "missing_state"), + [ + ({}, {}, "mx_server_url"), + ({"mx_server_url": "mx:8001"}, {"model": None}, "model"), + ( + {"mx_server_url": "mx:8001"}, + {"source_identity": None}, + "source_identity", + ), + ( + {"mx_server_url": "mx:8001"}, + {"model_config": None}, + "model_config", + ), + ], +) +def test_missing_mx_state_logs_reason_and_uses_native_hf_loader( + loader_kwargs, + load_overrides, + missing_state, +): + loader, weight_loader, _ = _loader(**loader_kwargs) + kwargs = _load_kwargs(**load_overrides) - @staticmethod - def _no_model(stack): # noqa: ARG004 - return MXCheckpointLoader(mx_server_url="http://mx:8001"), {} + with patch.object(checkpoint_loader_mod.logger, "info") as log_info: + value = loader.load_weights("checkpoint", **kwargs) - @staticmethod - def _upstream_raises(stack): - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress( - load_weights_side_effect=RuntimeError("boom"), - source_instances=[_source_instance(identity, post_transform=False)], - ) - stack.enter_context(_install_fake_modelexpress(fake_mx)) - return (loader, {"model": MagicMock(), "source_identity": identity}) - - @staticmethod - def _source_probe_raises(stack): - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress() - fake_mx.trtllm_live_transfer.MxClient.return_value.list_sources.side_effect = RuntimeError( - "server unavailable" - ) - stack.enter_context(_install_fake_modelexpress(fake_mx)) - return (loader, {"model": MagicMock(), "source_identity": identity}) - - @pytest.mark.parametrize( - "trigger_id, setup", - [ - ("no_mx_server_url", _no_url), - ("no_model_kwarg", _no_model), - ("source_probe_raises", _source_probe_raises), - ("upstream_raises", _upstream_raises), - ], - ids=[ - "no-mx-server-url", - "no-model-kwarg", - "source-probe-raises", - "upstream-raises", - ], + assert value == weight_loader.load_weights.return_value + log_info.assert_any_call( + f"MX loading unavailable: missing {missing_state}; " + "falling back to native Hugging Face checkpoint loading." + ) + weight_loader.load_weights.assert_called_once_with( + "checkpoint", + mapping=kwargs["mapping"], ) - 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.is_weights_preloaded() is False, ( - f"trigger={trigger_id}: is_weights_preloaded() must stay False on any fallback path" - ) - mock_super_load.assert_called_once() - - def test_missing_modelexpress_client_fails_with_install_hint(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - with ( - _block_modelexpress(), - patch.object(HfCheckpointLoader, "load_weights") as mock_super_load, - pytest.raises(ImportError) as exc_info, - ): - loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - ) - - message = str(exc_info.value) - assert 'pip install "tensorrt-llm[mx]"' in message - assert "select a different `checkpoint_format`" in message - assert loader.is_weights_preloaded() is False - mock_super_load.assert_not_called() - - -# --------------------------------------------------------------------------- -# load_weights: MX-success and mixed-success paths with 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 uses the empty dict plus is_weights_preloaded() - # to skip the standard weight-mapping pipeline. - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress( - load_weights_return={}, - source_instances=[_source_instance(identity, post_transform=False)], - ) - mapping = MagicMock(name="mapping") - model = MagicMock(name="model") - prepare_receiver = MagicMock() - - with _install_fake_modelexpress(fake_mx): - result = loader.load_weights( - "/nonexistent", - mapping=mapping, - model=model, - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result == {} - assert loader.is_weights_preloaded() is True - assert loader.is_post_transform_weights_preloaded() is False - prepare_receiver.assert_not_called() - - # 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. - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fallback_weight = MagicMock() - fallback_weight.numel.return_value = 2 - fallback_weight.element_size.return_value = 4 - fallback = {"some.weight": fallback_weight} - fake_mx = _build_fake_modelexpress( - load_weights_return=fallback, - source_instances=[_source_instance(identity, post_transform=False)], - ) - with ( - _install_fake_modelexpress(fake_mx), - patch.object(HfCheckpointLoader, "load_weights") as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - ) - - assert loader.is_weights_preloaded() is True - assert loader.is_post_transform_weights_preloaded() is False - assert result is fallback - mock_super_load.assert_not_called() - - def test_post_transform_full_success_prepares_receiver_before_p2p(self) -> None: - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - model = MagicMock(name="model") - events = [] - prepare_receiver = MagicMock(side_effect=lambda _model: events.append("prepare_receiver")) - fake_mx = _build_fake_modelexpress( - load_weights_side_effect=lambda *_args, **_kwargs: events.append("p2p") or {}, - source_instances=[_source_instance(identity)], +def test_missing_trtllm_adapter_uses_native_hf_loader(monkeypatch): + def fail_import(_name): + raise ModuleNotFoundError( + "No module named 'modelexpress.engines.trtllm'", + name="modelexpress.engines.trtllm", ) - with _install_fake_modelexpress(fake_mx): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=model, - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result == {} - assert loader.is_weights_preloaded() is True - assert loader.is_post_transform_weights_preloaded() is True - prepare_receiver.assert_called_once_with(model) - assert events == ["prepare_receiver", "p2p"] - - def test_post_transform_source_without_receiver_preparer_falls_back_before_p2p( - self, - ) -> None: - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - disk_weights = {"disk.weight": MagicMock()} - fake_mx = _build_fake_modelexpress( - load_weights_return={}, - source_instances=[_source_instance(identity)], - ) + monkeypatch.setattr(checkpoint_loader_mod, "import_module", fail_import) + loader, weight_loader, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs() - with ( - _install_fake_modelexpress(fake_mx), - patch.object( - HfCheckpointLoader, "load_weights", return_value=disk_weights - ) as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - allow_post_transform_weights=True, - ) - - assert result is disk_weights - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - mx_loader = fake_mx.trtllm_live_transfer.MxLiveWeightLoader.return_value - mx_loader.load_weights.assert_not_called() - mock_super_load.assert_called_once() - - def test_post_transform_source_falls_back_before_p2p_when_profile_is_not_qualified( - self, - ) -> None: - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - disk_weights = {"disk.weight": MagicMock()} - prepare_receiver = MagicMock() - fake_mx = _build_fake_modelexpress( - load_weights_return={}, - source_instances=[_source_instance(identity)], - ) + with patch.object(checkpoint_loader_mod.logger, "warning") as log_warning: + value = loader.load_weights("checkpoint", **kwargs) - with ( - _install_fake_modelexpress(fake_mx), - patch.object( - HfCheckpointLoader, "load_weights", return_value=disk_weights - ) as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - allow_post_transform_weights=False, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result is disk_weights - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - mx_loader = fake_mx.trtllm_live_transfer.MxLiveWeightLoader.return_value - mx_loader.load_weights.assert_not_called() - prepare_receiver.assert_not_called() - mock_super_load.assert_called_once() - - @pytest.mark.parametrize( - "protocol_value", - ["2", "not-an-int", _MISSING], - ids=["newer-protocol", "invalid-protocol", "missing-protocol"], + assert value == weight_loader.load_weights.return_value + log_warning.assert_any_call( + "The installed ModelExpress package does not provide the TensorRT-LLM " + "adapter; install modelexpress>=0.5.1 or another compatible version. " + "Falling back to native Hugging Face checkpoint loading." ) - def test_post_transform_source_with_unsupported_protocol_falls_back_before_p2p( - self, protocol_value - ) -> None: - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - disk_weights = {"disk.weight": MagicMock()} - prepare_receiver = MagicMock() - source_instance = _source_instance(identity) - if protocol_value is _MISSING: - source_instance.metadata.pop(_MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY) - else: - source_instance.metadata[_MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY] = protocol_value - fake_mx = _build_fake_modelexpress( - load_weights_return={}, - source_instances=[source_instance], - ) - - with ( - _install_fake_modelexpress(fake_mx), - patch.object( - HfCheckpointLoader, "load_weights", return_value=disk_weights - ) as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result is disk_weights - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - mx_loader = fake_mx.trtllm_live_transfer.MxLiveWeightLoader.return_value - mx_loader.load_weights.assert_not_called() - prepare_receiver.assert_not_called() - mock_super_load.assert_called_once() - - @pytest.mark.parametrize( - "transform_abi_id", - ["trtllm-llama-target-layout-v2", _MISSING], - ids=["mismatched-abi", "missing-abi"], + weight_loader.load_weights.assert_called_once_with( + "checkpoint", + mapping=kwargs["mapping"], ) - def test_post_transform_source_with_unsupported_abi_falls_back_before_p2p( - self, transform_abi_id: object - ) -> None: - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - disk_weights = {"disk.weight": MagicMock()} - prepare_receiver = MagicMock() - source_instance = _source_instance(identity) - if transform_abi_id is _MISSING: - source_instance.metadata.pop(_MX_TRANSFORM_ABI_ID_METADATA_KEY) - else: - source_instance.metadata[_MX_TRANSFORM_ABI_ID_METADATA_KEY] = transform_abi_id - fake_mx = _build_fake_modelexpress( - load_weights_return={}, - source_instances=[source_instance], - ) - - with ( - _install_fake_modelexpress(fake_mx), - patch.object( - HfCheckpointLoader, "load_weights", return_value=disk_weights - ) as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result is disk_weights - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - mx_loader = fake_mx.trtllm_live_transfer.MxLiveWeightLoader.return_value - mx_loader.load_weights.assert_not_called() - prepare_receiver.assert_not_called() - mock_super_load.assert_called_once() - - def test_selects_matching_source_metadata_from_multiple_instances(self) -> None: - rank0_identity = _identity(rank=0) - rank1_identity = _identity(rank=1) - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - model = MagicMock(name="model") - prepare_receiver = MagicMock() - fake_mx = _build_fake_modelexpress( - load_weights_return={}, - source_instances=[ - _source_instance(rank0_identity), - _source_instance(rank1_identity), - ], - ) - - with ( - _install_fake_modelexpress(fake_mx), - patch.object(HfCheckpointLoader, "load_weights") as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=model, - source_identity=rank1_identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result == {} - assert loader.is_weights_preloaded() is True - assert loader.is_post_transform_weights_preloaded() is True - prepare_receiver.assert_called_once_with(model) - mock_super_load.assert_not_called() - - def test_post_transform_mixed_success_falls_back_to_full_disk_load(self): - # Post-transform sources are safe only when all tensors arrive via P2P. - # If MX returns raw fallback tensors, avoid mixing layouts by falling - # back to a full disk load. - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fallback = {"some.weight": MagicMock(numel=lambda: 1, element_size=lambda: 4)} - disk_weights = {"disk.weight": MagicMock()} - model = MagicMock(name="model") - prepare_receiver = MagicMock() - fake_mx = _build_fake_modelexpress( - load_weights_return=fallback, - source_instances=[_source_instance(identity)], - ) - with ( - _install_fake_modelexpress(fake_mx), - patch.object( - HfCheckpointLoader, "load_weights", return_value=disk_weights - ) as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=model, - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result is disk_weights - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - prepare_receiver.assert_called_once_with(model) - mock_super_load.assert_called_once() - - def test_post_transform_source_can_be_disallowed_before_p2p(self): - # Some receiver shapes, such as target+draft speculative decoding, are - # not ready to mix post-transform target bytes with separately loaded - # raw draft bytes. Let ModelLoader force a disk fallback before MX - # starts RDMA, rather than accepting bytes it cannot safely stage. - identity = _identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - disk_weights = {"disk.weight": MagicMock()} - prepare_receiver = MagicMock() - fake_mx = _build_fake_modelexpress( - load_weights_return={}, - source_instances=[_source_instance(identity)], - ) - - with ( - _install_fake_modelexpress(fake_mx), - patch.object( - HfCheckpointLoader, "load_weights", return_value=disk_weights - ) as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - allow_post_transform_weights=False, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result is disk_weights - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - fake_mx.trtllm_live_transfer.MxLiveWeightLoader.assert_not_called() - prepare_receiver.assert_not_called() - mock_super_load.assert_called_once() - - -# --------------------------------------------------------------------------- -# 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(), source_identity=_identity()) # must not raise - - def test_publish_called_with_model(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - identity = _identity() - fake_mx = _build_fake_modelexpress() - model = MagicMock(name="model") - - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(model, source_identity=identity) - - fake_mx.trtllm_live_transfer.publish_model_params.assert_called_once() - args, kwargs = fake_mx.trtllm_live_transfer.publish_model_params.call_args - assert args == (model,) - metadata = kwargs["metadata"] - assert metadata[_MX_WEIGHT_LAYOUT_METADATA_KEY] == _MX_WEIGHT_LAYOUT_POST_TRANSFORM - assert metadata[_MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY] == str( - _MX_STAGED_TRANSFORM_PROTOCOL_VERSION - ) - assert metadata[_MX_TRANSFORM_ABI_ID_METADATA_KEY] == LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1 - assert _MX_SOURCE_IDENTITY_METADATA_KEY in metadata - - def test_publish_synchronizes_cuda_before_exposing_source(self, monkeypatch): - events = [] - monkeypatch.setattr( - mx_checkpoint_loader, - "_synchronize_cuda_for_mx_publish", - lambda: events.append("synchronize"), - ) - def _publish(*_args, **_kwargs): - events.append("publish") - - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress(publish_side_effect=_publish) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock(), source_identity=_identity()) - - assert events == ["synchronize", "publish"] - - def test_source_identity_required_for_post_transform_publish(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress() - - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock()) - - fake_mx.trtllm_live_transfer.publish_model_params.assert_not_called() - - def test_transform_abi_required_for_post_transform_publish(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - fake_mx = _build_fake_modelexpress() - - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source( - MagicMock(), - source_identity=_identity(transform_abi_id=None), - ) - - fake_mx.trtllm_live_transfer.publish_model_params.assert_not_called() - - def test_publish_without_metadata_kwarg_uses_identity_metadata(self): - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - calls = [] - captured = {} - - def _publish_without_metadata(model): - calls.append(model) - captured["identity"] = fake_mx.trtllm_live_transfer._build_trtllm_identity( - model_name="local-model" - ) - - fake_mx = _build_fake_modelexpress(publish_model_params=_publish_without_metadata) - source_identity = _identity() - with _install_fake_modelexpress(fake_mx): - model = MagicMock() - loader.publish_as_source(model, source_identity=source_identity) - - assert calls == [model] - metadata = captured["identity"].extra_parameters - serialized_identity = json.loads(metadata[_MX_SOURCE_IDENTITY_METADATA_KEY]) - assert "model_name" not in serialized_identity - assert SourceIdentity.from_dict(serialized_identity).matches(source_identity).matched - assert metadata[_MX_WEIGHT_LAYOUT_METADATA_KEY] == _MX_WEIGHT_LAYOUT_POST_TRANSFORM - assert metadata[_MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY] == str( - _MX_STAGED_TRANSFORM_PROTOCOL_VERSION - ) - assert metadata[_MX_TRANSFORM_ABI_ID_METADATA_KEY] == LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1 - - def test_env_var_set_during_publish_then_restored(self): - loader = MXCheckpointLoader(mx_server_url="http://mx-instance:9999") - captured_env = {} - - def _capture(model, **_kwargs): - 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(), source_identity=_identity()) - 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(), source_identity=_identity()) - 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(), source_identity=_identity()) # must not raise - - def test_publish_attaches_trtllm_source_identity_to_mx_identity(self): - source_identity = _source_identity() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001") - loader._local_source_identity = source_identity - captured = {} - - def _publish_side_effect(model, **_kwargs): - identity = fake_mx.trtllm_live_transfer._build_trtllm_identity(model_name="local-model") - captured["identity"] = identity - - fake_mx = _build_fake_modelexpress(publish_side_effect=_publish_side_effect) - - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source( - MagicMock(), - checkpoint_dir="/scratch/local-model", - source_identity=source_identity, - ) - - serialized = captured["identity"].extra_parameters["trtllm_source_identity"] - published_identity = SourceIdentity.from_dict(json.loads(serialized)) - assert published_identity.model_name is None - assert published_identity.matches(source_identity).matched - assert ( - captured["identity"].extra_parameters[_MX_WEIGHT_LAYOUT_METADATA_KEY] - == _MX_WEIGHT_LAYOUT_POST_TRANSFORM - ) - assert captured["identity"].extra_parameters[ - _MX_TRANSFORM_PROTOCOL_VERSION_METADATA_KEY - ] == str(_MX_STAGED_TRANSFORM_PROTOCOL_VERSION) - assert ( - captured["identity"].extra_parameters[_MX_TRANSFORM_ABI_ID_METADATA_KEY] - == LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1 - ) +def test_missing_modelexpress_package_logs_install_hint_and_uses_native_hf_loader( + monkeypatch, +): + def fail_import(_name): + raise ModuleNotFoundError("No module named 'modelexpress'", name="modelexpress") - def test_serialized_identity_ignores_local_checkpoint_path(self): - donor_identity = _identity() - receiver_payload = donor_identity.to_dict() - receiver_payload["model_name"] = "/tmp/no-shards/TinyLlama" - receiver_identity = SourceIdentity.from_dict(receiver_payload) + monkeypatch.setattr(checkpoint_loader_mod, "import_module", fail_import) + loader, weight_loader, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs() - assert donor_identity.model_name != receiver_identity.model_name - assert _serialize_source_identity(donor_identity) == _serialize_source_identity( - receiver_identity - ) + with patch.object(checkpoint_loader_mod.logger, "warning") as log_warning: + value = loader.load_weights("checkpoint", **kwargs) + assert value == weight_loader.load_weights.return_value + log_warning.assert_any_call( + "ModelExpress is not installed; install it with " + '`pip install "tensorrt-llm[mx]"`. Falling back to native ' + "Hugging Face checkpoint loading." + ) + weight_loader.load_weights.assert_called_once_with( + "checkpoint", + mapping=kwargs["mapping"], + ) -# --------------------------------------------------------------------------- -# Helpers — fake modelexpress modules and import blockers -# --------------------------------------------------------------------------- +def test_incompatible_trtllm_adapter_logs_missing_entrypoint(monkeypatch): + module = ModuleType("modelexpress.engines.trtllm") + monkeypatch.setitem(sys.modules, module.__name__, module) + loader, weight_loader, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs() -def _modelexpress_module_names(): - return [ - "modelexpress", - "modelexpress.trtllm_live_transfer", - ] + with patch.object(checkpoint_loader_mod.logger, "warning") as log_warning: + value = loader.load_weights("checkpoint", **kwargs) + assert value == weight_loader.load_weights.return_value + log_warning.assert_any_call( + "The installed ModelExpress TensorRT-LLM adapter is incompatible: " + "MxModelLoader is missing. Install a compatible ModelExpress version. " + "Falling back to native Hugging Face checkpoint loading." + ) + weight_loader.load_weights.assert_called_once_with( + "checkpoint", + mapping=kwargs["mapping"], + ) -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 +@pytest.mark.parametrize("fallback", ["missing-state", "missing-adapter"]) +def test_native_fallback_releases_previous_mx_session(monkeypatch, fallback): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=False, + value={"disk": object()}, + ) + loader, weight_loader, _ = _loader(mx_server_url="mx:8001") + loader.load_weights("checkpoint", **_load_kwargs()) + previous_session = instances[0] - 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 + kwargs = _load_kwargs() + if fallback == "missing-state": + kwargs["model"] = None + else: + monkeypatch.setitem(sys.modules, "modelexpress.engines.trtllm", ModuleType("empty")) - return _Blocker() + assert loader.load_weights("checkpoint", **kwargs) == weight_loader.load_weights.return_value + previous_session.cleanup.assert_called_once_with() + assert loader._mx_loader is None + loader.post_load_publish( + MagicMock(), + checkpoint_dir="checkpoint", + weights_preloaded=False, + ) + previous_session.publish_model.assert_not_called() -def _build_fake_modelexpress( - *, - load_weights_return=None, - load_weights_side_effect=None, - publish_side_effect=None, - publish_model_params=None, - source_instances=None, - source_metadata=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") +def test_auxiliary_native_load_preserves_previous_mx_session(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=True, + value={}, + ) + loader, weight_loader, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs() + loader.load_weights("checkpoint", **kwargs) + session = instances[0] + + draft_mapping = MagicMock() + value = loader.load_weights( + "draft-checkpoint", + mapping=draft_mapping, + _preserve_mx_session=True, + ) - # 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 []) - if source_metadata is not None: - client_instance.get_source_metadata.return_value = source_metadata - fake_trtllm_live.MxClient = MagicMock(return_value=client_instance) - fake_trtllm_live._build_trtllm_identity = MagicMock( - return_value=SimpleNamespace(extra_parameters={}) + assert value == weight_loader.load_weights.return_value + weight_loader.load_weights.assert_called_with( + "draft-checkpoint", + mapping=draft_mapping, ) + session.cleanup.assert_not_called() + assert loader._mx_loader is session + assert loader.is_weights_preloaded() + assert loader.is_post_transform_weights_preloaded() + + loader.post_load_publish( + kwargs["model"], + checkpoint_dir="checkpoint", + weights_preloaded=True, + source_identity=kwargs["source_identity"], + ) + session.publish_model.assert_called_once_with(kwargs["model"]) - # publish_model_params(model) - if publish_model_params is not None: - fake_trtllm_live.publish_model_params = publish_model_params - elif 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 test_trtllm_adapter_dependency_error_is_not_hidden(monkeypatch): + def fail_import(_name): + raise ModuleNotFoundError("No module named 'nixl'", name="nixl") + monkeypatch.setattr(checkpoint_loader_mod, "import_module", fail_import) + loader, _, _ = _loader(mx_server_url="mx:8001") -def _install_fake_modelexpress(fake_pkg): - """Context manager that installs a fake ``modelexpress`` into sys.modules.""" + with pytest.raises(ModuleNotFoundError, match="nixl"): + loader.load_weights("checkpoint", **_load_kwargs()) - 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 test_qualified_llama_delegates_to_shared_chain(monkeypatch): + value = {} + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=True, + value=value, + ) + loader, _, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs() + + assert loader.load_weights("checkpoint", **kwargs) is value + + session = instances[0] + assert session.model is kwargs["model"] + assert session.kwargs["checkpoint_loader"] is loader + assert session.kwargs["checkpoint_dir"] == "checkpoint" + assert session.kwargs["mapping"] is kwargs["mapping"] + assert session.kwargs["source_identity"] is kwargs["source_identity"] + assert session.kwargs["model_config"] is kwargs["model_config"] + assert session.kwargs["load_config"] is kwargs["load_config"] + assert session.kwargs["p2p_enabled"] is True + assert session.kwargs["transform_protocol_version"] == 1 + assert loader.is_weights_preloaded() + assert loader.is_post_transform_weights_preloaded() + + +def test_unqualified_model_keeps_rdma_unavailable(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=False, + value={"disk": object()}, + ) + loader, _, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs( + allow_post_transform_weights=False, + prepare_post_transform_receiver=None, + ) - 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 + loader.load_weights("checkpoint", **kwargs) - return _Installer() + assert instances[0].kwargs["p2p_enabled"] is False + assert not loader.is_weights_preloaded() + assert not loader.is_post_transform_weights_preloaded() -# --------------------------------------------------------------------------- -# Item 2: defensive MX_SOURCE_QUERY_TIMEOUT default -# --------------------------------------------------------------------------- +def test_qualified_model_requires_receiver_preparation(monkeypatch): + _install_fake_mx(monkeypatch, p2p_succeeded=True, value={}) + loader, _, _ = _loader(mx_server_url="mx:8001") + with pytest.raises( + RuntimeError, + match="requires receiver structure preparation", + ): + loader.load_weights( + "checkpoint", + **_load_kwargs(prepare_post_transform_receiver=None), + ) -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. Existing user values are preserved. - """ +def test_qualified_model_requires_transform_protocol(monkeypatch): + _install_fake_mx(monkeypatch, p2p_succeeded=True, value={}) + loader, _, _ = _loader(mx_server_url="mx:8001") - @pytest.fixture(autouse=True) - def _isolated_env(self, monkeypatch): - monkeypatch.delenv("MX_SOURCE_QUERY_TIMEOUT", raising=False) - yield + with pytest.raises( + RuntimeError, + match="requires a transform protocol version", + ): + loader.load_weights( + "checkpoint", + **_load_kwargs(post_transform_protocol_version=None), + ) - def test_no_registered_source_gets_short_default_during_load(self): - identity = _identity() - def _assert_timeout(*args, **kwargs): - assert os.environ.get("MX_SOURCE_QUERY_TIMEOUT") == "30" - return {} +@pytest.mark.parametrize( + "source_identity", + [ + replace(_source_identity(), format_version=SOURCE_IDENTITY_FORMAT_VERSION - 1), + replace(_source_identity(), transform_abi_id=None), + ], + ids=["old-format", "missing-transform-abi"], +) +def test_qualified_model_requires_current_source_identity(monkeypatch, source_identity): + _install_fake_mx(monkeypatch, p2p_succeeded=True, value={}) + loader, _, _ = _loader(mx_server_url="mx:8001") - 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(), - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=lambda _model: None, - ) - mx_loader = fake_mx.trtllm_live_transfer.MxLiveWeightLoader.return_value - mx_loader.load_weights.assert_called_once() - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - - def test_no_registered_source_honors_configured_timeout(self): - identity = _identity() - - def _assert_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_timeout) - with _install_fake_modelexpress(fake_mx): - loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=lambda _model: None, - ) - mx_loader = fake_mx.trtllm_live_transfer.MxLiveWeightLoader.return_value - mx_loader.load_weights.assert_called_once() - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - - def test_zero_timeout_falls_back_before_receiver_preparation(self): - identity = _identity() - disk_weights = {"disk.weight": MagicMock()} - prepare_receiver = MagicMock() - loader = MXCheckpointLoader(mx_server_url="http://mx:8001", query_timeout_s=0) - fake_mx = _build_fake_modelexpress() - - with ( - _install_fake_modelexpress(fake_mx), - patch.object( - HfCheckpointLoader, "load_weights", return_value=disk_weights - ) as mock_super_load, - ): - result = loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - allow_post_transform_weights=True, - prepare_post_transform_receiver=prepare_receiver, - ) - - assert result is disk_weights - assert loader.is_weights_preloaded() is False - assert loader.is_post_transform_weights_preloaded() is False - prepare_receiver.assert_not_called() - fake_mx.trtllm_live_transfer.MxLiveWeightLoader.assert_not_called() - mock_super_load.assert_called_once() - - def test_existing_source_keeps_upstream_default_when_unset(self): - identity = _identity() - - 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=[_source_instance(identity, post_transform=False)], - ) - with _install_fake_modelexpress(fake_mx): - loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - ) - assert "MX_SOURCE_QUERY_TIMEOUT" not in os.environ - - def test_env_value_preserved(self, monkeypatch): - identity = _identity() - # 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, - source_instances=[_source_instance(identity, post_transform=False)], - ) - with _install_fake_modelexpress(fake_mx): - loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - ) - assert os.environ.get("MX_SOURCE_QUERY_TIMEOUT") == "120" - - def test_configured_timeout_applies_during_load_and_restores_env(self): - identity = _identity() - - 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, - source_instances=[_source_instance(identity, post_transform=False)], + with pytest.raises( + RuntimeError, + match="current TRT-LLM SourceIdentity format and a transform-layout ABI", + ): + loader.load_weights( + "checkpoint", + **_load_kwargs(source_identity=source_identity), ) - with _install_fake_modelexpress(fake_mx): - loader.load_weights( - "/nonexistent", - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - ) - 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/Qwen3-8B") - assert loader.model_name == "Qwen/Qwen3-8B" - - 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/Qwen3-8B", "Qwen/Qwen3-8B"), - ("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"), - ], - ids=[ - "bare-name", - "hub-id", - "nested-hub-id", - "abs-path-simple", - "abs-path-nested", - "dot-relative", - "home-expansion", - "empty", - ], - ) - 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//" - # to "/" instead of the commit sha. - snapshot = "/cache/huggingface/hub/models--Qwen--Qwen3-8B/snapshots/abc123def456789" - assert _normalize_model_identity(snapshot) == "Qwen/Qwen3-8B" - 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" +def test_incompatible_transfer_protocol_fails_closed(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=True, + value={}, + transform_protocol_version=2, + ) + loader, _, _ = _loader(mx_server_url="mx:8001") + with pytest.raises( + RuntimeError, + match="compatible TRT-LLM transform protocol and SourceIdentity ABI", + ): + loader.load_weights("checkpoint", **_load_kwargs()) -class TestResolveMxModelName: - """_resolve_mx_model_name uses publish_as_source's lookup order.""" + assert not loader.is_weights_preloaded() + assert not loader.is_post_transform_weights_preloaded() + instances[0].cleanup.assert_called_once_with() + assert loader._mx_loader is None - @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_p2p_transfer_cannot_return_native_weights(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=True, + value={"disk": object()}, + ) + loader, _, _ = _loader(mx_server_url="mx:8001") - 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" + with pytest.raises( + RuntimeError, + match="MX P2P loading must not return native checkpoint weights", + ): + loader.load_weights("checkpoint", **_load_kwargs()) - def test_basename_fallback_when_arg_and_env_missing(self): - assert _resolve_mx_model_name(None, "/scratch/local-model") == "local-model" + assert not loader.is_weights_preloaded() + assert not loader.is_post_transform_weights_preloaded() + instances[0].cleanup.assert_called_once_with() + assert loader._mx_loader is None - def test_snapshot_fallback_when_arg_and_env_missing(self): - snapshot = "/cache/huggingface/hub/models--Qwen--Qwen3-8B/snapshots/abc123" - assert _resolve_mx_model_name(None, snapshot) == "Qwen/Qwen3-8B" - def test_unknown_when_all_missing(self): - assert _resolve_mx_model_name(None, None) == "unknown" - assert _resolve_mx_model_name("", None) == "unknown" +def test_failed_mx_load_cleans_session_and_prevents_publish(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=False, + value=None, + load_error=RuntimeError("transfer failed"), + ) + loader, _, _ = _loader(mx_server_url="mx:8001") - 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" + with pytest.raises(RuntimeError, match="transfer failed"): + loader.load_weights("checkpoint", **_load_kwargs()) + session = instances[0] + session.cleanup.assert_called_once_with() + assert loader._mx_loader is None -class TestLoadWeightsModelName: - def test_uses_resolved_model_name_during_load_and_restores_env(self, monkeypatch): - monkeypatch.setenv("MODEL_NAME", "prior-model") - identity = _identity() - snapshot = "/cache/hub/models--Other--Model/snapshots/abc123" + loader.post_load_publish( + MagicMock(), + checkpoint_dir="checkpoint", + weights_preloaded=False, + ) + session.publish_model.assert_not_called() - def _assert_model_name(*args, **kwargs): - assert os.environ.get("MODEL_NAME") == "Qwen/Qwen3-8B" - return {} - loader = MXCheckpointLoader( - mx_server_url="http://mx:8001", - model_name="Qwen/Qwen3-8B", - ) - fake_mx = _build_fake_modelexpress( - load_weights_side_effect=_assert_model_name, - source_instances=[_source_instance(identity, post_transform=False)], - ) +def test_repeated_load_clears_cleaned_session_before_replacement(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=False, + value={"disk": object()}, + ) + loader, _, _ = _loader(mx_server_url="mx:8001") + loader.load_weights("checkpoint", **_load_kwargs()) + first_session = instances[0] - with _install_fake_modelexpress(fake_mx): - loader.load_weights( - snapshot, - mapping=MagicMock(), - model=MagicMock(), - source_identity=identity, - ) - - assert os.environ.get("MODEL_NAME") == "prior-model" - fake_mx.trtllm_live_transfer._build_trtllm_identity.assert_called_with( - model_name="Qwen/Qwen3-8B" - ) + class FailingMxModelLoader: + def __init__(self, **_kwargs): + raise RuntimeError("construction failed") + sys.modules["modelexpress.engines.trtllm"].MxModelLoader = FailingMxModelLoader -class TestPublishAsSourceModelName: - """publish_as_source sets MODEL_NAME for upstream publish_model_params. + with pytest.raises(RuntimeError, match="construction failed"): + loader.load_weights("checkpoint", **_load_kwargs()) - The prior environment value must be restored afterwards. - """ + first_session.cleanup.assert_called_once_with() + assert loader._mx_loader is None - @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 + loader.cleanup() + first_session.cleanup.assert_called_once_with() - def test_uses_explicit_constructor_model_name(self): - loader = MXCheckpointLoader( - mx_server_url="http://mx:8001", - model_name="Qwen/Qwen3-8B", - ) - captured = {} - - def _capture(model, **_kwargs): - 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(), source_identity=_identity()) - - assert captured["MODEL_NAME"] == "Qwen/Qwen3-8B" - 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, **_kwargs): - 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", - source_identity=_identity(), - ) - - 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--Qwen3-8B/snapshots/abc123def456" - captured = {} - - def _capture(model, **_kwargs): - 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, - source_identity=_identity(), - ) - - # Critical: NOT the commit hash, the human-readable Hub-ID form. - assert captured["MODEL_NAME"] == "Qwen/Qwen3-8B" - 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, **_kwargs): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") +def test_p2p_receiver_republishes_after_trt_post_load(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=True, + value={}, + ) + loader, _, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs() + model = kwargs["model"] + loader.load_weights("checkpoint", **kwargs) + instances[0].publish_model.assert_not_called() + + loader.post_load_publish( + model, + checkpoint_dir="checkpoint", + weights_preloaded=True, + source_identity=kwargs["source_identity"], + ) - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock(), source_identity=_identity()) + instances[0].publish_model.assert_called_once_with(model) - 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 test_native_source_publishes_after_trt_post_load(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=False, + value={"disk": object()}, + ) + loader, _, _ = _loader(mx_server_url="mx:8001") + kwargs = _load_kwargs() + model = kwargs["model"] + loader.load_weights("checkpoint", **kwargs) + instances[0].publish_model.assert_not_called() + + loader.post_load_publish( + model, + checkpoint_dir="checkpoint", + weights_preloaded=False, + source_identity=kwargs["source_identity"], + ) - def _capture(model, **_kwargs): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") + instances[0].publish_model.assert_called_once_with(model) - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock(), source_identity=_identity()) - assert captured["MODEL_NAME"] == "from-env-only" - assert os.environ.get("MODEL_NAME") == "from-env-only" +def test_cleanup_releases_mx_and_native_loader_resources(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=False, + value={"disk": object()}, + ) + loader, weight_loader, config_loader = _loader(mx_server_url="mx:8001") + loader.load_weights("checkpoint", **_load_kwargs()) - 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 = {} + loader.cleanup() - def _capture(model, **_kwargs): - captured["MODEL_NAME"] = os.environ.get("MODEL_NAME") + instances[0].cleanup.assert_called_once_with() + weight_loader.cleanup.assert_called_once_with() + config_loader.cleanup.assert_called_once_with() + assert loader._mx_loader is None - fake_mx = _build_fake_modelexpress(publish_side_effect=_capture) - with _install_fake_modelexpress(fake_mx): - loader.publish_as_source(MagicMock(), source_identity=_identity()) - assert captured["MODEL_NAME"] == "unknown" +def test_cleanup_continues_when_mx_cleanup_fails(monkeypatch): + instances = _install_fake_mx( + monkeypatch, + p2p_succeeded=False, + value={"disk": object()}, + ) + loader, weight_loader, config_loader = _loader(mx_server_url="mx:8001") + loader.load_weights("checkpoint", **_load_kwargs()) + cleanup_error = RuntimeError("cleanup failed") + instances[0].cleanup.side_effect = cleanup_error + warning = MagicMock() + monkeypatch.setattr(checkpoint_loader_mod.logger, "warning", warning) + + loader.cleanup() + + instances[0].cleanup.assert_called_once_with() + weight_loader.cleanup.assert_called_once_with() + config_loader.cleanup.assert_called_once_with() + assert loader._mx_loader is None + warning.assert_called_once_with( + f"Failed to clean up ModelExpress loader {instances[0]!r}: {cleanup_error!r}" + ) diff --git a/tests/unittest/_torch/weight_sharing/test_mx_source_identity_gate.py b/tests/unittest/_torch/weight_sharing/test_mx_source_identity_gate.py index a5cf9c87c717..93d5b3d34e6c 100644 --- a/tests/unittest/_torch/weight_sharing/test_mx_source_identity_gate.py +++ b/tests/unittest/_torch/weight_sharing/test_mx_source_identity_gate.py @@ -1,127 +1,102 @@ # 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 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""Exercise the real MX checkpoint loader pre-transfer SourceIdentity gate. - -Upstream `modelexpress` is never imported. The tests run the compatibility -decision against real `SourceIdentity` objects and use a small discovery -client stub for the pinned ModelExpress 0.4.1 API shape. -""" +"""Exercise the real ModelExpress identity gate before RDMA mutation.""" +import json +from dataclasses import replace from types import SimpleNamespace +from unittest.mock import MagicMock import pytest -from _source_identity_fakes import FakeMapping -from _source_identity_fakes import make_identity as _identity +from _source_identity_fakes import make_identity -from tensorrt_llm._torch.models.checkpoints.mx.checkpoint_loader import ( - MXCheckpointLoader, - _build_mx_source_metadata, +from tensorrt_llm._torch.weight_sharing import ( + LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1, + SOURCE_IDENTITY_FORMAT_VERSION, ) -pytestmark = pytest.mark.cpu_only +pytest.importorskip("modelexpress") +import modelexpress.load_strategy as load_strategy # noqa: E402 +from modelexpress.engines import trtllm # noqa: E402 +from modelexpress.load_strategy import rdma_strategy # noqa: E402 -def _new_loader(local_identity): - """Construct a loader while bypassing the heavy base initializer.""" - loader = MXCheckpointLoader.__new__(MXCheckpointLoader) - loader._local_source_identity = local_identity - return loader - - -def test_gate_proceeds_on_matching_identity(): - local = _identity(attn_backend="TRTLLM") - source = _identity(attn_backend="TRTLLM") - loader = _new_loader(local) - assert loader._source_metadata_identity_compatible(_build_mx_source_metadata(source)) is True +def _identity(*, transform_abi_id=LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1): + return replace( + make_identity(model_name="meta-llama/Llama-3.1-8B-Instruct"), + transform_abi_id=transform_abi_id, + ) -def test_gate_falls_back_on_mismatch(): - local = _identity(attn_backend="TRTLLM") - source = _identity(attn_backend="FLASHINFER") - loader = _new_loader(local) - assert loader._source_metadata_identity_compatible(_build_mx_source_metadata(source)) is False +def _mx_modules(): + return trtllm, load_strategy, rdma_strategy -def test_gate_falls_back_on_checkpoint_artifact_mismatch(): - local = _identity(artifact_key="fine-tune-a") - source = _identity(artifact_key="fine-tune-b") - loader = _new_loader(local) - assert loader._source_metadata_identity_compatible(_build_mx_source_metadata(source)) is False +def test_real_adapter_serializes_authoritative_trt_identity(): + trtllm, _, _ = _mx_modules() + identity = _identity() -def test_gate_falls_back_when_no_local_identity(): - # MX must not consume shared weights unless the receiver identity exists. - loader = _new_loader(None) - assert ( - loader._source_metadata_identity_compatible(_build_mx_source_metadata(_identity())) is False + mx_identity = trtllm.build_mx_identity( + identity, + transform_protocol_version=1, + ) + serialized = json.loads(mx_identity.extra_parameters["trtllm_source_identity"]) + + assert serialized == { + key: value for key, value in identity.to_dict().items() if key != "model_name" + } + assert serialized["format_version"] == SOURCE_IDENTITY_FORMAT_VERSION + assert serialized["transform_abi_id"] == LLAMA_POST_TRANSFORM_LAYOUT_ABI_V1 + + +@pytest.mark.parametrize( + "incompatible_identity", + [ + _identity(transform_abi_id="trtllm-llama-target-layout-v2"), + replace(_identity(), format_version=SOURCE_IDENTITY_FORMAT_VERSION - 1), + replace(_identity(), backend_fingerprint="different-backend"), + replace( + _identity(), + artifact_identity=replace(_identity().artifact_identity, digest="1" * 64), + ), + ], + ids=["transform-abi", "format", "backend", "artifact"], +) +def test_incompatible_identity_falls_back_before_receiver_mutation( + incompatible_identity, +): + trtllm, load_strategy, rdma_strategy = _mx_modules() + target_identity = trtllm.build_mx_identity( + _identity(), + transform_protocol_version=1, + ) + published_identity = trtllm.build_mx_identity( + incompatible_identity, + transform_protocol_version=1, ) - - -def test_gate_falls_back_when_source_identity_unavailable(): - loader = _new_loader(_identity()) - assert loader._source_metadata_identity_compatible(None) is False - - -def test_fetch_source_metadata_supports_modelexpress_0_4_1_client_shape_and_close_failure(): - local = _identity() - loader = MXCheckpointLoader.__new__(MXCheckpointLoader) - loader._local_source_identity = local - loader._mx_server_url = "http://mx:8001" - loader._model_name = None class _Client: - def __init__(self, *, server_url): - self.server_url = server_url - - def get_metadata(self, mx_source_id, worker_id): - raise AssertionError("ID-based metadata lookup should not be used for identity queries") - - def list_sources(self, *, identity): - return SimpleNamespace( - instances=[SimpleNamespace(mx_source_id="source", worker_id="worker")] - ) - - def close(self): - raise RuntimeError("close failed") - - metadata = loader._fetch_source_metadata( - "ckpt", - _Client, - lambda **_kw: SimpleNamespace(extra_parameters={}), + def __init__(self): + self.queried_identity = None + + def list_sources(self, *, identity, status_filter): + self.queried_identity = identity + instances = [MagicMock()] if identity == published_identity else [] + return SimpleNamespace(instances=instances) + + client = _Client() + ctx = SimpleNamespace( + mx_client=client, + identity=target_identity, + global_rank=0, + worker_rank=0, ) + result = load_strategy.LoadResult(value=MagicMock(), model=MagicMock()) - assert metadata == _build_mx_source_metadata(local) + with pytest.raises(load_strategy.StrategyFailed) as error: + rdma_strategy.RdmaStrategy().load(result, ctx) - -def test_load_weights_pops_source_identity_kwarg(): - # source_identity must be consumed by load_weights and never leak into the - # HfCheckpointLoader disk-fallback signature. With no server URL configured - # the loader takes the disk-fallback path; we only assert the kwarg is - # stored and not forwarded. - loader = MXCheckpointLoader.__new__(MXCheckpointLoader) - loader._mx_server_url = None - loader._p2p_succeeded = False - captured = {} - - def fake_fallback(checkpoint_dir, mapping, *, reason=None, **kwargs): - captured["kwargs"] = kwargs - return {} - - loader._fallback_to_disk = fake_fallback - identity = _identity() - loader.load_weights("ckpt", mapping=FakeMapping(), model=None, source_identity=identity) - assert loader._local_source_identity is identity - assert "source_identity" not in captured["kwargs"] - assert "model" not in captured["kwargs"] + assert error.value.mutated is False + assert client.queried_identity == target_identity + assert target_identity != published_identity diff --git a/tests/unittest/llmapi/test_mx_args.py b/tests/unittest/llmapi/test_mx_args.py index 058d6073f69d..db6a5eea2312 100644 --- a/tests/unittest/llmapi/test_mx_args.py +++ b/tests/unittest/llmapi/test_mx_args.py @@ -12,7 +12,7 @@ import pytest -from tensorrt_llm.llmapi.llm_args import TorchLlmArgs +from tensorrt_llm.llmapi.llm_args import ModelExpressConfig, TorchLlmArgs pytestmark = pytest.mark.cpu_only @@ -59,6 +59,10 @@ def test_mx_server_query_timeout_accepts_nonnegative_int(self): args = _make_args(checkpoint_format="MX", mx_config={"server_query_timeout_s": 1200}) assert args.mx_config.server_query_timeout_s == 1200 + def test_mx_server_query_timeout_is_deprecated(self): + field = ModelExpressConfig.model_json_schema()["properties"]["server_query_timeout_s"] + assert field["deprecated"] is True + def test_mx_server_query_timeout_rejects_negative(self): with pytest.raises(ValueError): _make_args(checkpoint_format="MX", mx_config={"server_query_timeout_s": -1})