diff --git a/examples/p2p_transfer_k8s/client/trtllm/Dockerfile b/examples/p2p_transfer_k8s/client/trtllm/Dockerfile new file mode 100644 index 000000000..6140933e5 --- /dev/null +++ b/examples/p2p_transfer_k8s/client/trtllm/Dockerfile @@ -0,0 +1,26 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +ARG TRTLLM_IMAGE=nvcr.io/nvidia/tensorrt-llm/release:1.3.0rc11@sha256:d91c80ba8baf763782b1078267ed6b1e06363bebff4961094bf6e5679d371d04 +FROM ${TRTLLM_IMAGE} + +USER root + +ENV NIXL_PREFIX=/opt/nvidia/nvda_nixl \ + NIXL_LIB_DIR=/opt/nvidia/nvda_nixl/lib/x86_64-linux-gnu \ + NIXL_PLUGIN_DIR=/opt/nvidia/nvda_nixl/lib/x86_64-linux-gnu/plugins +ENV LD_LIBRARY_PATH=${NIXL_LIB_DIR}:${NIXL_PLUGIN_DIR}:/usr/local/tensorrt/targets/x86_64-linux-gnu/lib:${LD_LIBRARY_PATH} + +COPY examples/p2p_transfer_k8s/client/trtllm/install_modelexpress_client.sh /tmp/trtllm_modelexpress/install_modelexpress_client.sh +COPY examples/p2p_transfer_k8s/client/trtllm/fix_nixl_runpath.py /tmp/trtllm_modelexpress/fix_nixl_runpath.py + +COPY modelexpress_client/python /tmp/modelexpress_client +COPY modelexpress_common/proto /tmp/proto + +RUN bash /tmp/trtllm_modelexpress/install_modelexpress_client.sh && \ + python3 /tmp/trtllm_modelexpress/fix_nixl_runpath.py + +RUN python3 -c "import modelexpress; print('ModelExpress OK')" && \ + test -f "$NIXL_PLUGIN_DIR/libplugin_UCX.so" && \ + echo "TRT-LLM ModelExpress image OK" && \ + rm -rf /tmp/trtllm_modelexpress /tmp/modelexpress_client /tmp/proto diff --git a/examples/p2p_transfer_k8s/client/trtllm/Dockerfile.ph3-gcp-gb200 b/examples/p2p_transfer_k8s/client/trtllm/Dockerfile.ph3-gcp-gb200 deleted file mode 100644 index c64c5b76d..000000000 --- a/examples/p2p_transfer_k8s/client/trtllm/Dockerfile.ph3-gcp-gb200 +++ /dev/null @@ -1,78 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -# Phase 3 P2P image for GCP GB200 (ARM64) -# Layers ModelExpress client + Dynamo P2P hooks + PRESHARDED patches -# on karenc's Dynamo v1.0.0 image -# -# Base: nvcr.io/nvidian/dynamo-dev/karenc:dynamo-trtllm-v1.0.0-a9b6f95 -# - TRT-LLM 1.3.0rc5 -# - NIXL pre-installed -# - ARM64 (aarch64) -# - Dynamo worker (python3 -m dynamo.trtllm) -# -# Requires dynamo fork at ../dynamo (the Dynamo branch with P2P support) -# -# Build (from modelexpress repo root): -# docker buildx build --platform linux/arm64 --no-cache \ -# -f examples/p2p_transfer_trtllm/Dockerfile.ph3-gcp-gb200 \ -# --build-context dynamo=../dynamo \ -# -t /dynamo-trtllm-mx: \ -# --push . -FROM nvcr.io/nvidian/dynamo-dev/karenc:dynamo-trtllm-v1.0.0-a9b6f95 - -USER root - -# Multinode MPI: sshd requires privilege separation directory owned by root, mode 0755 -RUN mkdir -p /run/sshd && chmod 0755 /run/sshd && chown root:root /run/sshd - -# Install ModelExpress client (gRPC + NIXL transfer) -COPY modelexpress_client/python /tmp/modelexpress_client -COPY modelexpress_common/proto /tmp/proto - -RUN cd /tmp/modelexpress_client && \ - pip install --no-cache-dir grpcio grpcio-tools protobuf && \ - python3 -m grpc_tools.protoc \ - -I/tmp/proto \ - --python_out=modelexpress \ - --grpc_python_out=modelexpress \ - /tmp/proto/p2p.proto && \ - sed -i 's/^import p2p_pb2/from . import p2p_pb2/' modelexpress/p2p_pb2_grpc.py && \ - pip install --no-cache-dir . && \ - rm -rf /tmp/modelexpress_client /tmp/proto - -# Apply Dynamo P2P hooks (--model-express-url support) -# From the Dynamo branch with P2P support: backend_args.py, engine.py, llm_worker.py -COPY --from=dynamo components/src/dynamo/trtllm/backend_args.py /tmp/dynamo_p2p/backend_args.py -COPY --from=dynamo components/src/dynamo/trtllm/engine.py /tmp/dynamo_p2p/engine.py -COPY --from=dynamo components/src/dynamo/trtllm/workers/llm_worker.py /tmp/dynamo_p2p/llm_worker.py - -RUN DYNAMO_PKG=/opt/dynamo/venv/lib/python3.12/site-packages/dynamo/trtllm && \ - cp /tmp/dynamo_p2p/backend_args.py "$DYNAMO_PKG/backend_args.py" && \ - cp /tmp/dynamo_p2p/engine.py "$DYNAMO_PKG/engine.py" && \ - cp /tmp/dynamo_p2p/llm_worker.py "$DYNAMO_PKG/workers/llm_worker.py" && \ - rm -rf /tmp/dynamo_p2p && \ - grep -q "model.express.url\|model_express_url" "$DYNAMO_PKG/backend_args.py" && \ - echo "Dynamo P2P hooks OK" - -# Apply PRESHARDED patches to TRT-LLM 1.3.0rc5 -COPY trtllm_patches/v1.3.0rc5/apply_patches.py /tmp/apply_patches.py -RUN python3 /tmp/apply_patches.py && rm /tmp/apply_patches.py - -# Patch tp_allgather to use chunked allgather (fixes MPI_ERR_TRUNCATE with ob1 TCP BTL) -COPY trtllm_patches/v1.3.0rc5/patch_tp_allgather.py /tmp/patch_tp_allgather.py -RUN python3 /tmp/patch_tp_allgather.py && rm /tmp/patch_tp_allgather.py - -# Patch model_loader: source publishes BEFORE post_load_weights, target runs full post_load_weights -COPY trtllm_patches/v1.3.0rc5/patch_model_loader.py /tmp/patch_model_loader.py -RUN python3 /tmp/patch_model_loader.py && rm /tmp/patch_model_loader.py - -# Verify installation (can't import full TRT-LLM at build time — no GPU) -RUN python3 -c "from modelexpress.client import MxClient; print('ModelExpress client OK')" && \ - python3 -c "from modelexpress.trtllm_live_transfer import publish_model_params; print('publish_model_params OK')" && \ - grep -q "PRESHARDED" /opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/llmapi/llm_args.py && \ - echo "LoadFormat.PRESHARDED OK" && \ - grep -q "publish_model_params" /opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/_torch/pyexecutor/model_loader.py && \ - echo "model_loader.py source publish hook OK" && \ - grep -q "MODEL_EXPRESS_SOURCE" /opt/dynamo/venv/lib/python3.12/site-packages/dynamo/trtllm/engine.py && \ - echo "engine.py MODEL_EXPRESS_SOURCE OK" diff --git a/examples/p2p_transfer_k8s/client/trtllm/README.md b/examples/p2p_transfer_k8s/client/trtllm/README.md index f1fa675fe..b2c823556 100644 --- a/examples/p2p_transfer_k8s/client/trtllm/README.md +++ b/examples/p2p_transfer_k8s/client/trtllm/README.md @@ -1,301 +1,36 @@ -# ModelExpress P2P for TRT-LLM on GCP GB200 +# TensorRT-LLM ModelExpress Test Image -GPU-to-GPU weight loading for Kimi K2.5 via ModelExpress NIXL RDMA. -Targets load weights in seconds instead of 15-24 minutes from PVC. +This directory contains only the minimum files needed to build a temporary +TRT-LLM image for ModelExpress checkpoint-loader validation. -## Architecture +The image starts from a TRT-LLM release image, installs the local ModelExpress +Python client, and expects the base TRT-LLM image to already include the +ModelExpress checkpoint-loader hooks. -``` - ┌─────────────────────────────────────┐ - │ ModelExpress Server + Redis │ - │ (gRPC metadata, NIXL descriptors) │ - └────────┬────────────────┬────────────┘ - publish │ │ query - metadata │ │ metadata - │ │ - ┌──────────────────────────┴───┐ ┌─────────┴───────────────────────┐ - │ MX Source (DGD, TP=8) │ │ MX Targets (DGD) │ - │ │ │ │ - │ 1. Load weights from PVC │ │ 1. Query MX server for source │ - │ 2. model.load_weights() │ │ 2. NIXL RDMA into param bufs │ - │ 3. ── PUBLISH HERE ── │ │ 3. post_load_weights() │ - │ 4. post_load_weights() │ │ 4. NCCL init + serve │ - │ 5. Serve (holds GPU mem) │ │ │ - │ │ │ Source publishes BEFORE step 4 │ - │ Node A ┌──┐┌──┐┌──┐┌──┐ │ │ so targets run same xforms │ - │ │R0││R1││R2││R3│ │ │ │ - │ Node B ┌──┐┌──┐┌──┐┌──┐ │ │ ┌───────────────────────────┐ │ - │ │R4││R5││R6││R7│ │ │ │ Prefill (TP=8, 2 nodes) │ │ - │ │ │ │ Decode (TP=8, 2 nodes) │ │ - └──────────────────────────────┘ │ │ Frontend (KV router) │ │ - │ │ └───────────────────────────┘ │ - │ NIXL RDMA │ │ - └────────────────────► 400-457 Gbps RoCE per rank │ - 90.75 GB/rank ─────────────────────────────────┘ -``` +The Dockerfile pins the default TRT-LLM base image to +`nvcr.io/nvidia/tensorrt-llm/release:1.3.0rc11@sha256:d91c80ba8baf763782b1078267ed6b1e06363bebff4961094bf6e5679d371d04` +for reproducible validation builds. Override `TRTLLM_IMAGE` when testing a +TRT-LLM image that includes the ModelExpress checkpoint-loader hooks. ---- +## Build -## Quick Start - -### Prerequisites +Run from the ModelExpress repo root: ```bash -# 1. Teleport auth -tsh kube login dynamo-gcp-dev-01 - -# 2. Dynamo platform (etcd + NATS) must be running -kubectl -n default get pods # verify etcd-0, nats-0 - -# 3. Secrets -kubectl -n default get secret hf-token-secret -kubectl -n default get secret nvcr-imagepullsecret - -# 4. ComputeDomain (creates IMEX channels for GPU allocation) -cat < | grep "published ALL" -``` - -### Step 3: Deploy target (receives weights via P2P) - -```bash -# After source publishes: -kubectl -n default apply -f kimi-target-agg-tp8-dgd.yaml -# TP=8, loads via RDMA in ~2 seconds -# Watch: kubectl -n default logs -f | grep "Gbps" -``` - -### Step 4: Test inference - -```bash -FRONTEND=$(kubectl -n default get pod -l app.kubernetes.io/part-of=kimi-target-agg-tp8 \ - -l nvidia.com/dynamo-component=frontend -o name | head -1) -kubectl -n default exec $FRONTEND -- curl -s http://localhost:8000/v1/chat/completions \ - -H "Content-Type: application/json" \ - -d '{"model":"baseten-admin/Kimi-2.5-text-nvfp4-v3", - "messages":[{"role":"user","content":"What is the capital of France?"}], - "max_tokens":50}' -``` - ---- - -## Disaggregated Inference (P2P) - -Separate prefill and decode workers with KV cache transfer. Each loads -weights via P2P from the same MX source. - -### Step 1: Deploy infrastructure + source - -```bash -kubectl -n default apply -f mx-infra-decode.yaml -kubectl -n default apply -f kimi-source-decode-dgd.yaml -# Wait for source to publish (~15 min) -``` - -### Step 2: Deploy disagg targets - -```bash -kubectl -n default apply -f kimi-disagg-mx-tp8-dgd.yaml -# Creates: Frontend (KV router) + Prefill (MX target) + Decode (MX target) -# Both load via P2P concurrently -``` - -### Step 3: Test inference - -```bash -FRONTEND_IP=$(kubectl -n default get pod -l app.kubernetes.io/part-of=kimi-disagg-mx-tp8 \ - -l nvidia.com/dynamo-component=frontend -o jsonpath='{.items[0].status.podIP}') -kubectl -n default exec -- curl -s http://${FRONTEND_IP}:8000/v1/chat/completions \ - -H "Content-Type: application/json" \ - -d '{"model":"baseten-admin/Kimi-2.5-text-nvfp4-v3", - "messages":[{"role":"user","content":"What is the capital of France?"}], - "max_tokens":50}' -``` - -### Disagg without MX (PVC baseline) - -To validate disagg works without P2P (both workers load from PVC): - -```bash -kubectl -n default apply -f kimi-disagg-baseline-dgd.yaml -# No MX source needed — loads directly from shared-model-cache PVC -``` - ---- - -## File Reference - -### Infrastructure - -| File | Purpose | -|------|---------| -| `mx-infra.yaml` | ModelExpress server + Redis (TP=4 source) | -| `mx-infra-decode.yaml` | ModelExpress server + Redis (TP=8 source) | - -### Sources (load from PVC, publish for RDMA) - -| File | TP | Nodes | Notes | -|------|-----|-------|-------| -| `kimi-source-decode-dgd.yaml` | 8 | 2 | Primary source for TP=8 targets | -| `kimi-source-deploy.yaml` | 4 | 1 | Legacy TP=4 source (Deployment) | -| `kimi-source-dgd.yaml` | 4 | 1 | TP=4 source (DGD format) | - -### Targets (receive weights via P2P) - -| File | Mode | TP | MX Source | Notes | -|------|------|-----|-----------|-------| -| `kimi-target-agg-tp8-dgd.yaml` | Aggregated | 8 | decode server | Simplest P2P test | -| `kimi-disagg-mx-tp8-dgd.yaml` | Disagg | 8+8 | decode server | Prefill + decode via P2P | -| `kimi-disagg-baseline-dgd.yaml` | Disagg | 8+8 | PVC (no MX) | Baseline for comparison | -| `kimi-disagg-mx-dgd.yaml` | Disagg | 4+4 | MX server | Legacy TP=4 disagg | -| `kimi-disagg-phase2-dgd.yaml` | Disagg | 4+8 | dual MX | Mixed TP (needs 2 sources) | -| `kimi-target-pvc-tp8-dgd.yaml` | Aggregated | 8 | PVC (no MX) | Ground truth baseline | - -### Testing - -| File | Purpose | -|------|---------| -| `qwen-source-deploy.yaml` | Qwen 0.5B TP=2 source (fast iteration) | -| `qwen-target-deploy.yaml` | Qwen 0.5B TP=2 target | -| `qwen-baseline-test.yaml` | Qwen baseline without MX | - ---- - -## Required Configuration for GCP GB200 - -All worker pods need these settings. See `docs/disagg_trtllm.md` for details. - -```yaml -securityContext: - privileged: true # RDMA memory registration - runAsUser: 0 # SSH key path fix for multinode - -env: - HOME: /root # Fix SSH host key path mismatch - UCX_TLS: "self,sm,rc,cuda_copy,gdr_copy,tcp" - UCX_IB_GID_INDEX: "3" # GCP RoCEv2 GID selection - TRTLLM_UCX_INTERFACE: eth0 # Prevent 169.254.x.x binding - OMPI_MCA_pml: ob1 # Bypass UCX UD for MPI - OMPI_MCA_btl: "tcp,self,vader" - NATS_SERVER: "nats://dynamo-platform-nats..svc.cluster.local:4222" - ETCD_ENDPOINTS: "dynamo-platform-etcd..svc.cluster.local:2379" - -# Engine config (extra-engine-args YAML): - cache_transceiver_config: - backend: DEFAULT # Bypass UCX UD for KV cache transfer - enable_autotuner: false # Avoid warmup MPI desync -``` - ---- - -## Cleanup - -```bash -# Delete all workloads -kubectl -n default delete dgd --all - -# Delete compute domain (releases IMEX channels) -kubectl -n default delete computedomain my-compute-domain - -# Keep infra running for next deployment -# kubectl -n default delete -f mx-infra-decode.yaml # if needed -``` - ---- - -## Building the Image - -The image combines three repos into a single ARM64 container layered on the -Dynamo TRT-LLM base image. - -### Repos and branches - -| Repo | Branch | What it provides | -|------|--------|-----------------| -| `modelexpress` | your modelexpress branch | MX client, NIXL transfer, TRT-LLM patches | -| `dynamo` | the Dynamo branch with P2P support | Engine P2P hooks (`--model-express-url`) | -| `TensorRT-LLM` | your TRT-LLM branch | `LoadFormat.PRESHARDED` (applied via patches) | - -### Directory layout - -``` -~/work/github/ -├── modelexpress/ (your modelexpress branch) -└── dynamo/ (the Dynamo branch with P2P support) -``` - -### Build command - -```bash -cd ~/work/github/modelexpress - -docker buildx build --platform linux/arm64 --no-cache \ - -f examples/p2p_transfer_trtllm/Dockerfile.ph3-gcp-gb200 \ - --build-context dynamo=../dynamo \ - -t /dynamo-trtllm-mx: \ - --push . -``` - -The Dockerfile (`examples/p2p_transfer_trtllm/Dockerfile.ph3-gcp-gb200`): -1. Starts from `karenc:dynamo-trtllm-v1.0.0-a9b6f95` (TRT-LLM 1.3.0rc5 + NIXL, ARM64) -2. Installs ModelExpress Python client (gRPC + NIXL transfer) -3. Copies Dynamo engine/worker files from `dynamo` repo via `--build-context` -4. Applies TRT-LLM patches: `PRESHARDED` LoadFormat, source publish hook, MPI allgather fix - -### Building the base image from dynamo - -If you don't have access to `karenc:dynamo-trtllm-v1.0.0-a9b6f95`, build the -base image from the `dynamo` repo using its rendered Dockerfile: - -```bash -cd ~/work/github/dynamo - -docker buildx build --platform linux/arm64 --no-cache \ - -f container/trtllm-runtime-cuda13.1-arm64-rendered.Dockerfile \ - --build-arg ARCH=arm64 \ - --build-arg ARCH_ALT=aarch64 \ - -t my-registry/dynamo-trtllm-base:latest \ +docker buildx build --platform linux/amd64 \ + -f examples/p2p_transfer_k8s/client/trtllm/Dockerfile \ + --build-arg TRTLLM_IMAGE=nvcr.io/nvidia/tensorrt-llm/release: \ + -t /trtllm-modelexpress: \ --push . ``` -Then update the `FROM` line in `Dockerfile.ph3-gcp-gb200` to point to your -base image instead of `karenc:dynamo-trtllm-v1.0.0-a9b6f95`. +## Files -### Current image +| File | Purpose | +| --- | --- | +| `Dockerfile` | Builds the temporary TRT-LLM + ModelExpress e2e image. | +| `install_modelexpress_client.sh` | Builds protobuf stubs and installs the local ModelExpress Python client. | +| `fix_nixl_runpath.py` | Makes the NIXL Python binding resolve TRT-LLM's system NIXL libraries. | -``` -/dynamo-trtllm-mx: -``` +No Kubernetes manifests are kept here yet. The current TRT-LLM integration is +still under validation, and stale DGD manifests were intentionally removed. diff --git a/examples/p2p_transfer_k8s/client/trtllm/fix_nixl_runpath.py b/examples/p2p_transfer_k8s/client/trtllm/fix_nixl_runpath.py new file mode 100644 index 000000000..309114877 --- /dev/null +++ b/examples/p2p_transfer_k8s/client/trtllm/fix_nixl_runpath.py @@ -0,0 +1,67 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Make the NIXL Python binding use TRT-LLM's system NIXL stack.""" + +from __future__ import annotations + +import glob +import os +import site +import subprocess +import sys +import sysconfig + + +def _run(args: list[str]) -> str: + return subprocess.run(args, check=True, text=True, capture_output=True).stdout + + +def _site_package_dirs() -> list[str]: + candidates = [] + candidates.extend(site.getsitepackages()) + purelib = sysconfig.get_paths().get("purelib") + if purelib: + candidates.append(purelib) + candidates.extend(sys.path) + + dirs = [] + seen = set() + for candidate in candidates: + if not candidate: + continue + path = os.path.abspath(candidate) + if path in seen or not os.path.isdir(path): + continue + seen.add(path) + dirs.append(path) + return dirs + + +def main() -> None: + nixl_lib_dir = os.environ["NIXL_LIB_DIR"] + bindings = [] + for site_pkg in _site_package_dirs(): + bindings.extend( + glob.glob(os.path.join(site_pkg, "nixl_cu13", "_bindings*.so")) + ) + if len(bindings) != 1: + raise RuntimeError(f"Expected one nixl_cu13 binding, found {bindings}") + + binding = bindings[0] + rpath = "$ORIGIN/../.nixl_cu13.mesonpy.libs:$ORIGIN/../nixl_cu13.libs" + _run(["patchelf", "--set-rpath", rpath, binding]) + + dynamic = _run(["readelf", "-d", binding]) + if "RUNPATH" not in dynamic: + raise RuntimeError(f"{binding} does not use DT_RUNPATH after patching") + + linked = _run(["ldd", binding]) + expected = f"{nixl_lib_dir}/libnixl.so" + if expected not in linked: + raise RuntimeError(f"{binding} does not resolve libnixl.so from {nixl_lib_dir}") + + +if __name__ == "__main__": + main() diff --git a/examples/p2p_transfer_k8s/client/trtllm/install_modelexpress_client.sh b/examples/p2p_transfer_k8s/client/trtllm/install_modelexpress_client.sh new file mode 100644 index 000000000..6f1fbef03 --- /dev/null +++ b/examples/p2p_transfer_k8s/client/trtllm/install_modelexpress_client.sh @@ -0,0 +1,28 @@ +#!/usr/bin/env bash +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +set -euo pipefail + +cd /tmp/modelexpress_client + +pip install --no-cache-dir \ + "grpcio==1.66.2" \ + "grpcio-tools==1.66.2" \ + "protobuf>=5.27.0,<6.0.0" + +python3 -m grpc_tools.protoc \ + -I/tmp/proto \ + --python_out=modelexpress \ + --grpc_python_out=modelexpress \ + /tmp/proto/p2p.proto +sed -i 's/^import p2p_pb2/from . import p2p_pb2/' modelexpress/p2p_pb2_grpc.py + +pip install --no-cache-dir . + +# TRT-LLM 1.3.x release images are CUDA 13, so use the matching NIXL CUDA +# plugin instead of the package default pulled in by the generic dependency. +pip uninstall -y nixl nixl-cu12 nixl-cu13 +pip install --no-cache-dir --no-deps --force-reinstall \ + "nixl==0.10.1" \ + "nixl-cu13==0.10.1" diff --git a/examples/p2p_transfer_k8s/client/trtllm/kimi-disagg-mx-tp8-dgd.yaml b/examples/p2p_transfer_k8s/client/trtllm/kimi-disagg-mx-tp8-dgd.yaml deleted file mode 100644 index 2d714b14f..000000000 --- a/examples/p2p_transfer_k8s/client/trtllm/kimi-disagg-mx-tp8-dgd.yaml +++ /dev/null @@ -1,568 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -# Kimi K2.5 disaggregated with ModelExpress P2P weight loading -# Both prefill and decode pull weights via RDMA from MX source (TP=8) -# -# Config: prefill TP=8 (2 nodes) + decode TP=8 (2 nodes) = 4 nodes total -# Source: kimi-source-decode DGD publishes to modelexpress-server-decode -# -# Prerequisites: -# - MX source deployed and published (kimi-source-decode-dgd.yaml) -# - modelexpress-server-decode + redis-decode running (mx-infra-decode.yaml) -# - shared-model-cache PVC (for tokenizer/config only) -# - Dynamo platform (etcd + NATS) -# - ComputeDomain my-compute-domain with >= 4 nodes -# -# Usage: -# kubectl -n default apply -f kimi-disagg-mx-tp8-dgd.yaml ---- -apiVersion: v1 -kind: ConfigMap -metadata: - name: kimi-disagg-mx-tp8-config -data: - prefill.yaml: | - cache_transceiver_config: - backend: DEFAULT - max_tokens_in_buffer: 240000 - max_num_tokens: 4096 - enable_chunked_prefill: true - cuda_graph_config: - max_batch_size: 4 - enable_padding: true - disable_overlap_scheduler: true - enable_attention_dp: true - enable_autotuner: false - kv_cache_config: - dtype: fp8 - enable_block_reuse: true - free_gpu_memory_fraction: 0.7 - tokens_per_block: 32 - max_batch_size: 4 - max_seq_len: 234102 - model_kwargs: - num_hidden_layers: 61 - moe_config: - backend: TRTLLM - moe_expert_parallel_size: 8 - print_iter_log: true - tensor_parallel_size: 8 - trust_remote_code: true - decode.yaml: | - cache_transceiver_config: - backend: DEFAULT - max_tokens_in_buffer: 240000 - max_num_tokens: 8192 - cuda_graph_config: - max_batch_size: 8 - enable_padding: true - disable_overlap_scheduler: true - enable_attention_dp: true - enable_autotuner: false - kv_cache_config: - dtype: fp8 - enable_block_reuse: false - free_gpu_memory_fraction: 0.7 - tokens_per_block: 32 - max_batch_size: 8 - max_seq_len: 240000 - model_kwargs: - num_hidden_layers: 61 - moe_config: - backend: TRTLLM - use_low_precision_moe_combine: true - moe_expert_parallel_size: 8 - num_postprocess_workers: 4 - print_iter_log: true - stream_interval: 10 - tensor_parallel_size: 8 - trust_remote_code: true ---- -apiVersion: nvidia.com/v1alpha1 -kind: DynamoGraphDeployment -metadata: - name: kimi-disagg-mx-tp8 -spec: - services: - Frontend: - componentType: frontend - replicas: 1 - extraPodSpec: - affinity: - # Update node pool names for your cluster - nodeAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - nodeSelectorTerms: - - matchExpressions: - - key: cloud.google.com/gke-nodepool - operator: In - values: - - - containers: null - imagePullSecrets: - - name: nvcr-imagepullsecret - mainContainer: - image: "/dynamo-trtllm-mx:" - command: - - python3 - args: - - -m - - dynamo.frontend - - --router-mode - - kv - - --router-reset-states - - --request-plane - - nats - env: - - name: POD_UID - valueFrom: - fieldRef: - fieldPath: metadata.uid - - name: HF_HOME - value: /model-cache - - name: NATS_SERVER - value: "nats://dynamo-platform-nats.default.svc.cluster.local:4222" - - name: ETCD_ENDPOINTS - value: "dynamo-platform-etcd.default.svc.cluster.local:2379" - volumeMounts: - - mountPath: /model-cache - name: model-cache - nodeSelector: - kubernetes.io/arch: arm64 - nvidia.com/gpu.product: NVIDIA-GB200 - tolerations: - - effect: NoSchedule - key: dedicated - operator: Equal - value: user-workload - - effect: NoExecute - key: dedicated - operator: Equal - value: user-workload - - effect: NoSchedule - key: nvidia.com/gpu - operator: Exists - - effect: NoSchedule - key: kubernetes.io/arch - operator: Exists - volumes: - - name: model-cache - persistentVolumeClaim: - claimName: shared-model-cache - - prefill: - componentType: worker - subComponentType: prefill - envFromSecret: hf-token-secret - replicas: 1 - multinode: - nodeCount: 2 - annotations: - networking.gke.io/default-interface: eth0 - networking.gke.io/interfaces: | - [ - {"interfaceName":"eth0","network":"default"}, - {"interfaceName":"rdma0","network":"rdma-0"}, - {"interfaceName":"rdma1","network":"rdma-1"}, - {"interfaceName":"rdma2","network":"rdma-2"}, - {"interfaceName":"rdma3","network":"rdma-3"} - ] - resources: - limits: - gpu: "4" - custom: - networking.gke.io.networks/rdma-0: "1" - networking.gke.io.networks/rdma-1: "1" - networking.gke.io.networks/rdma-2: "1" - networking.gke.io.networks/rdma-3: "1" - claims: - - name: compute-domain-channel - readinessProbe: - httpGet: - path: /health - port: 9090 - periodSeconds: 10 - timeoutSeconds: 30 - failureThreshold: 60 - extraPodSpec: - securityContext: {} - affinity: - # Update node pool names for your cluster - nodeAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - nodeSelectorTerms: - - matchExpressions: - - key: cloud.google.com/gke-nodepool - operator: In - values: - - - podAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - - labelSelector: - matchExpressions: - - key: app.kubernetes.io/part-of - operator: In - values: - - kimi-disagg-mx-tp8 - topologyKey: nvidia.com/gpu.clique - containers: null - imagePullSecrets: - - name: nvcr-imagepullsecret - mainContainer: - image: "/dynamo-trtllm-mx:" - workingDir: /workspace/ - securityContext: - privileged: true - runAsUser: 0 - capabilities: - add: - - IPC_LOCK - command: - - python3 - - -m - - dynamo.trtllm - args: - - --model-path - - baseten-admin/Kimi-2.5-text-nvfp4-v3 - - --served-model-name - - baseten-admin/Kimi-2.5-text-nvfp4-v3 - - --extra-engine-args - - /config/prefill.yaml - - --disaggregation-mode - - prefill - - --model-express-url - - modelexpress-server-decode.default.svc.cluster.local:8001 - - --model-express-role - - target - - --tensor-parallel-size - - "8" - - --publish-events-and-metrics - - --request-plane - - nats - - --kv-block-size - - "32" - env: - - name: HOME - value: /root - - name: POD_UID - valueFrom: - fieldRef: - fieldPath: metadata.uid - - name: HF_HOME - value: /model-cache - - name: HF_HUB_OFFLINE - value: "1" - - name: TRITON_CACHE_DIR - value: /tmp/.triton-cache - - name: HF_MODULES_CACHE - value: /tmp/hf_modules - - name: MODEL_EXPRESS_URL - value: "modelexpress-server-decode.default.svc.cluster.local:8001" - - name: MODEL_NAME - value: "baseten-admin/Kimi-2.5-text-nvfp4-v3" - - name: NATS_SERVER - value: "nats://dynamo-platform-nats.default.svc.cluster.local:4222" - - name: ETCD_ENDPOINTS - value: "dynamo-platform-etcd.default.svc.cluster.local:2379" - - name: NCCL_CUMEM_ENABLE - value: "1" - - name: NCCL_NVLS_ENABLE - value: "1" - - name: NVIDIA_GDRCOPY - value: "1" - - name: NCCL_SOCKET_IFNAME - value: eth0 - - name: GLOO_SOCKET_IFNAME - value: eth0 - - name: NCCL_STORE_TIMEOUT - value: "7200" - - name: OMPI_MCA_pml - value: "ob1" - - name: OMPI_MCA_btl - value: "tcp,self,vader" - - name: OMPI_MCA_btl_tcp_if_include - value: "eth0" - - name: OMPI_MCA_oob_tcp_if_include - value: "eth0" - # Per-rank IB NIC pinning. Workaround for openucx/ucx#11259. - - name: MX_RDMA_NIC_PIN - value: "auto" - - name: UCX_TLS - value: "cuda_ipc,cuda_copy,rc" - - name: UCX_IB_GID_INDEX - value: "3" - - name: UCX_RC_TIMEOUT - value: "600s" - - name: UCX_KEEPALIVE_INTERVAL - value: "300s" - - name: UCX_LOG_LEVEL - value: debug - - name: NIXL_LOG_LEVEL - value: DEBUG - - name: TRTLLM_UCX_INTERFACE - value: eth0 - - name: TRTLLM_MOE_ENABLE_ALLTOALL_WITHOUT_ALLGATHER - value: "1" - - name: TRTLLM_ENABLE_PDL - value: "1" - - name: TRTLLM_SERVER_DISABLE_GC - value: "1" - - name: TRTLLM_WORKER_DISABLE_GC - value: "1" - - name: NCCL_GRAPH_MIXING_SUPPORT - value: "0" - - name: NCCL_MAX_NCHANNELS - value: "2" - - name: TLLM_LOG_LEVEL - value: "INFO" - startupProbe: - failureThreshold: 60 - httpGet: - path: /live - port: 9090 - periodSeconds: 60 - timeoutSeconds: 5 - volumeMounts: - - mountPath: /model-cache - name: shared-model-cache - - mountPath: /config - name: trtllm-config - readOnly: true - nodeSelector: - kubernetes.io/arch: arm64 - nvidia.com/gpu.product: NVIDIA-GB200 - tolerations: - - effect: NoSchedule - key: dedicated - operator: Equal - value: user-workload - - effect: NoExecute - key: dedicated - operator: Equal - value: user-workload - - effect: NoSchedule - key: nvidia.com/gpu - operator: Exists - - effect: NoSchedule - key: kubernetes.io/arch - operator: Exists - resourceClaims: - - name: compute-domain-channel - resourceClaimTemplateName: my-compute-domain-channel - volumes: - - name: shared-model-cache - persistentVolumeClaim: - claimName: shared-model-cache - - name: trtllm-config - configMap: - name: kimi-disagg-mx-tp8-config - - decode: - componentType: worker - subComponentType: decode - envFromSecret: hf-token-secret - replicas: 1 - multinode: - nodeCount: 2 - annotations: - networking.gke.io/default-interface: eth0 - networking.gke.io/interfaces: | - [ - {"interfaceName":"eth0","network":"default"}, - {"interfaceName":"rdma0","network":"rdma-0"}, - {"interfaceName":"rdma1","network":"rdma-1"}, - {"interfaceName":"rdma2","network":"rdma-2"}, - {"interfaceName":"rdma3","network":"rdma-3"} - ] - resources: - limits: - gpu: "4" - custom: - networking.gke.io.networks/rdma-0: "1" - networking.gke.io.networks/rdma-1: "1" - networking.gke.io.networks/rdma-2: "1" - networking.gke.io.networks/rdma-3: "1" - claims: - - name: compute-domain-channel - readinessProbe: - httpGet: - path: /health - port: 9090 - periodSeconds: 10 - timeoutSeconds: 30 - failureThreshold: 60 - extraPodSpec: - securityContext: {} - affinity: - # Update node pool names for your cluster - nodeAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - nodeSelectorTerms: - - matchExpressions: - - key: cloud.google.com/gke-nodepool - operator: In - values: - - - podAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - - labelSelector: - matchExpressions: - - key: app.kubernetes.io/part-of - operator: In - values: - - kimi-disagg-mx-tp8 - topologyKey: nvidia.com/gpu.clique - containers: null - imagePullSecrets: - - name: nvcr-imagepullsecret - mainContainer: - image: "/dynamo-trtllm-mx:" - workingDir: /workspace/ - securityContext: - privileged: true - runAsUser: 0 - capabilities: - add: - - IPC_LOCK - command: - - python3 - - -m - - dynamo.trtllm - args: - - --model-path - - baseten-admin/Kimi-2.5-text-nvfp4-v3 - - --served-model-name - - baseten-admin/Kimi-2.5-text-nvfp4-v3 - - --extra-engine-args - - /config/decode.yaml - - --disaggregation-mode - - decode - - --model-express-url - - modelexpress-server-decode.default.svc.cluster.local:8001 - - --model-express-role - - target - - --tensor-parallel-size - - "8" - - --publish-events-and-metrics - - --request-plane - - nats - - --kv-block-size - - "32" - env: - - name: HOME - value: /root - - name: POD_UID - valueFrom: - fieldRef: - fieldPath: metadata.uid - - name: HF_HOME - value: /model-cache - - name: HF_HUB_OFFLINE - value: "1" - - name: TRITON_CACHE_DIR - value: /tmp/.triton-cache - - name: HF_MODULES_CACHE - value: /tmp/hf_modules - - name: MODEL_EXPRESS_URL - value: "modelexpress-server-decode.default.svc.cluster.local:8001" - - name: MODEL_NAME - value: "baseten-admin/Kimi-2.5-text-nvfp4-v3" - - name: NATS_SERVER - value: "nats://dynamo-platform-nats.default.svc.cluster.local:4222" - - name: ETCD_ENDPOINTS - value: "dynamo-platform-etcd.default.svc.cluster.local:2379" - - name: NCCL_CUMEM_ENABLE - value: "1" - - name: NCCL_NVLS_ENABLE - value: "1" - - name: NVIDIA_GDRCOPY - value: "1" - - name: NCCL_SOCKET_IFNAME - value: eth0 - - name: GLOO_SOCKET_IFNAME - value: eth0 - - name: NCCL_STORE_TIMEOUT - value: "7200" - - name: OMPI_MCA_pml - value: "ob1" - - name: OMPI_MCA_btl - value: "tcp,self,vader" - - name: OMPI_MCA_btl_tcp_if_include - value: "eth0" - - name: OMPI_MCA_oob_tcp_if_include - value: "eth0" - # Per-rank IB NIC pinning. Workaround for openucx/ucx#11259. - - name: MX_RDMA_NIC_PIN - value: "auto" - - name: UCX_TLS - value: "cuda_ipc,cuda_copy,rc" - - name: UCX_IB_GID_INDEX - value: "3" - - name: UCX_RC_TIMEOUT - value: "600s" - - name: UCX_KEEPALIVE_INTERVAL - value: "300s" - - name: UCX_LOG_LEVEL - value: debug - - name: NIXL_LOG_LEVEL - value: DEBUG - - name: TRTLLM_UCX_INTERFACE - value: eth0 - - name: TRTLLM_MOE_ENABLE_ALLTOALL_WITHOUT_ALLGATHER - value: "1" - - name: TRTLLM_ENABLE_PDL - value: "1" - - name: TRTLLM_SERVER_DISABLE_GC - value: "1" - - name: TRTLLM_WORKER_DISABLE_GC - value: "1" - - name: NCCL_GRAPH_MIXING_SUPPORT - value: "0" - - name: NCCL_MAX_NCHANNELS - value: "2" - - name: ENABLE_CONFIGURABLE_MOE - value: "1" - - name: TLLM_LOG_LEVEL - value: "INFO" - startupProbe: - failureThreshold: 60 - httpGet: - path: /live - port: 9090 - periodSeconds: 60 - timeoutSeconds: 5 - volumeMounts: - - mountPath: /model-cache - name: shared-model-cache - - mountPath: /config - name: trtllm-config - readOnly: true - nodeSelector: - kubernetes.io/arch: arm64 - nvidia.com/gpu.product: NVIDIA-GB200 - tolerations: - - effect: NoSchedule - key: dedicated - operator: Equal - value: user-workload - - effect: NoExecute - key: dedicated - operator: Equal - value: user-workload - - effect: NoSchedule - key: nvidia.com/gpu - operator: Exists - - effect: NoSchedule - key: kubernetes.io/arch - operator: Exists - resourceClaims: - - name: compute-domain-channel - resourceClaimTemplateName: my-compute-domain-channel - volumes: - - name: shared-model-cache - persistentVolumeClaim: - claimName: shared-model-cache - - name: trtllm-config - configMap: - name: kimi-disagg-mx-tp8-config diff --git a/examples/p2p_transfer_k8s/client/trtllm/kimi-source-decode-dgd.yaml b/examples/p2p_transfer_k8s/client/trtllm/kimi-source-decode-dgd.yaml deleted file mode 100644 index 612f16c42..000000000 --- a/examples/p2p_transfer_k8s/client/trtllm/kimi-source-decode-dgd.yaml +++ /dev/null @@ -1,229 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -# Kimi K2.5 decode source — DGD with multinode TP=8 (2 nodes × 4 GPUs) -# Publishes weights via NIXL to modelexpress-server-decode -# Phase 2: provides TP=8 weights for decode targets -# Namespace: default -# -# Prerequisites: -# - modelexpress-server-decode + redis-decode running (mx-infra-decode.yaml) -# - shared-model-cache PVC with Kimi model -# - Dynamo platform (etcd + NATS) -# - DGD operator webhook running ---- -apiVersion: v1 -kind: ConfigMap -metadata: - name: kimi-source-decode-trtllm-config -data: - source-decode.yaml: | - max_num_tokens: 8192 - enable_chunked_prefill: true - disable_overlap_scheduler: true - cuda_graph_config: - max_batch_size: 8 - enable_padding: true - enable_attention_dp: true - kv_cache_config: - dtype: fp8 - enable_block_reuse: false - free_gpu_memory_fraction: 0.75 - tokens_per_block: 32 - max_batch_size: 8 - max_seq_len: 256000 - model_kwargs: - num_hidden_layers: 61 - moe_config: - backend: TRTLLM - use_low_precision_moe_combine: true - moe_expert_parallel_size: 8 - num_postprocess_workers: 4 - print_iter_log: true - stream_interval: 10 - tensor_parallel_size: 8 - trust_remote_code: true ---- -apiVersion: nvidia.com/v1alpha1 -kind: DynamoGraphDeployment -metadata: - name: kimi-source-decode -spec: - services: - source: - componentType: worker - envFromSecret: hf-token-secret - multinode: - nodeCount: 2 - replicas: 1 - resources: - limits: - gpu: "4" - claims: - - name: compute-domain-channel - extraPodSpec: - hostNetwork: true - dnsPolicy: ClusterFirstWithHostNet - securityContext: {} - affinity: - # Update node pool names for your cluster - nodeAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - nodeSelectorTerms: - - matchExpressions: - - key: cloud.google.com/gke-nodepool - operator: In - values: - - - podAffinity: - requiredDuringSchedulingIgnoredDuringExecution: - - labelSelector: - matchExpressions: - - key: app.kubernetes.io/part-of - operator: In - values: - - kimi-source-decode - topologyKey: nvidia.com/gpu.clique - containers: null - imagePullSecrets: - - name: nvcr-imagepullsecret - mainContainer: - image: "/dynamo-trtllm-mx:" - imagePullPolicy: Always - workingDir: /workspace/ - securityContext: - privileged: true - command: - - python3 - - -m - - dynamo.trtllm - args: - - --model-path - - baseten-admin/Kimi-2.5-text-nvfp4-v3 - - --served-model-name - - baseten-admin/Kimi-2.5-text-nvfp4-v3 - - --extra-engine-args - - /config/source-decode.yaml - - --model-express-url - - modelexpress-server-decode.default.svc.cluster.local:8001 - - --model-express-role - - source - - --tensor-parallel-size - - "8" - env: - - name: POD_UID - valueFrom: - fieldRef: - fieldPath: metadata.uid - - name: NATS_SERVER - value: "nats://dynamo-platform-nats.default.svc.cluster.local:4222" - - name: ETCD_ENDPOINTS - value: "dynamo-platform-etcd.default.svc.cluster.local:2379" - - name: MODEL_EXPRESS_URL - value: "modelexpress-server-decode.default.svc.cluster.local:8001" - - name: MODEL_NAME - value: "baseten-admin/Kimi-2.5-text-nvfp4-v3" - - name: MODEL_EXPRESS_SOURCE - value: "1" - - name: WORLD_SIZE - value: "8" - - name: HOME - value: "/root" - - name: HF_HOME - value: /model-cache - - name: HF_HUB_OFFLINE - value: "1" - - name: HF_MODULES_CACHE - value: /tmp/hf_modules - - name: TRITON_CACHE_DIR - value: /model-cache/.triton-cache - - name: NIXL_LOG_LEVEL - value: "INFO" - - name: TRTLLM_UCX_INTERFACE - value: "eth0" - # Per-rank IB NIC pinning. Workaround for openucx/ucx#11259. - - name: MX_RDMA_NIC_PIN - value: "auto" - - name: UCX_TLS - value: "self,sm,rc,cuda_copy,gdr_copy,tcp" - - name: UCX_IB_GID_INDEX - value: "3" - - name: OMPI_MCA_pml - value: "ob1" - - name: OMPI_MCA_btl - value: "tcp,self,vader" - - name: OMPI_MCA_btl_tcp_if_include - value: "eth0" - - name: OMPI_MCA_oob_tcp_if_include - value: "eth0" - - name: NCCL_NET_PLUGIN - value: "none" - - name: NCCL_CUMEM_ENABLE - value: "1" - - name: NCCL_NVLS_ENABLE - value: "1" - - name: NVIDIA_GDRCOPY - value: "1" - - name: NCCL_SOCKET_IFNAME - value: eth0 - - name: GLOO_SOCKET_IFNAME - value: eth0 - - name: NCCL_STORE_TIMEOUT - value: "7200" - - name: TRTLLM_MOE_ENABLE_ALLTOALL_WITHOUT_ALLGATHER - value: "1" - - name: TRTLLM_ENABLE_PDL - value: "1" - - name: TRTLLM_SERVER_DISABLE_GC - value: "1" - - name: TRTLLM_WORKER_DISABLE_GC - value: "1" - - name: NCCL_GRAPH_MIXING_SUPPORT - value: "0" - - name: NCCL_MAX_NCHANNELS - value: "2" - - name: ENABLE_CONFIGURABLE_MOE - value: "1" - - name: TLLM_LOG_LEVEL - value: "INFO" - - name: TLLM_OVERRIDE_LAYER_NUM - value: "61" - volumeMounts: - - mountPath: /dev/infiniband - name: infiniband - - mountPath: /model-cache - name: model-cache - - mountPath: /config - name: trtllm-config - readOnly: true - nodeSelector: - kubernetes.io/arch: arm64 - tolerations: - - effect: NoSchedule - key: dedicated - operator: Equal - value: user-workload - - effect: NoExecute - key: dedicated - operator: Equal - value: user-workload - - effect: NoSchedule - key: nvidia.com/gpu - operator: Exists - - effect: NoSchedule - key: kubernetes.io/arch - operator: Exists - resourceClaims: - - name: compute-domain-channel - resourceClaimTemplateName: my-compute-domain-channel - volumes: - - name: infiniband - hostPath: - path: /dev/infiniband - type: Directory - - name: model-cache - persistentVolumeClaim: - claimName: shared-model-cache - - name: trtllm-config - configMap: - name: kimi-source-decode-trtllm-config diff --git a/examples/p2p_transfer_k8s/client/trtllm/mx-infra-decode.yaml b/examples/p2p_transfer_k8s/client/trtllm/mx-infra-decode.yaml deleted file mode 100644 index ae27f85d2..000000000 --- a/examples/p2p_transfer_k8s/client/trtllm/mx-infra-decode.yaml +++ /dev/null @@ -1,104 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -# ModelExpress server + Redis for decode source (Phase 2 mixed TP) -# Separate from prefill MX server so each source publishes to its own store -# Deploys on nodes (amd64) -# Namespace: default ---- -apiVersion: apps/v1 -kind: Deployment -metadata: - name: redis-decode - labels: - app: redis-decode -spec: - replicas: 1 - selector: - matchLabels: - app: redis-decode - template: - metadata: - labels: - app: redis-decode - spec: - # Update node pool names for your cluster - nodeSelector: - cloud.google.com/gke-nodepool: - tolerations: - - key: dedicated - operator: Equal - value: user-workload - effect: NoExecute - containers: - - name: redis - image: redis:7-alpine - ports: - - containerPort: 6379 - resources: - requests: - cpu: "0.5" - memory: "512Mi" ---- -apiVersion: v1 -kind: Service -metadata: - name: redis-decode -spec: - selector: - app: redis-decode - ports: - - port: 6379 - targetPort: 6379 ---- -apiVersion: apps/v1 -kind: Deployment -metadata: - name: modelexpress-server-decode - labels: - app: modelexpress-server-decode -spec: - replicas: 1 - selector: - matchLabels: - app: modelexpress-server-decode - template: - metadata: - labels: - app: modelexpress-server-decode - spec: - # Update node pool names for your cluster - nodeSelector: - cloud.google.com/gke-nodepool: - tolerations: - - key: dedicated - operator: Equal - value: user-workload - effect: NoExecute - containers: - - name: server - image: nvcr.io/nvidia/ai-dynamo/modelexpress-server:0.3.0 - ports: - - containerPort: 8001 - env: - - name: MX_METADATA_BACKEND - value: "redis" - - name: REDIS_URL - value: "redis://redis-decode:6379" - resources: - requests: - cpu: "1" - memory: "1Gi" - imagePullSecrets: - - name: nvcr-imagepullsecret ---- -apiVersion: v1 -kind: Service -metadata: - name: modelexpress-server-decode -spec: - selector: - app: modelexpress-server-decode - ports: - - port: 8001 - targetPort: 8001 diff --git a/modelexpress_client/python/modelexpress/engines/sglang/adapter.py b/modelexpress_client/python/modelexpress/engines/sglang/adapter.py index 0cf143cc9..470a2a308 100644 --- a/modelexpress_client/python/modelexpress/engines/sglang/adapter.py +++ b/modelexpress_client/python/modelexpress/engines/sglang/adapter.py @@ -16,8 +16,16 @@ from ... import p2p_pb2 from ...adapter import EngineAdapter -from ...load_strategy.context import LoadContext, LoadResult -from ...metadata.client_factory import create_metadata_client +from ...load_strategy.context import ( + LoadContext, + LoadResult, + resolve_model_streamer_uri, +) +from ...metadata.client_factory import ( + create_metadata_client, + resolve_metadata_port, + resolve_metadata_server_url, +) logger = logging.getLogger("modelexpress.engines.sglang.adapter") @@ -40,6 +48,9 @@ def __init__( self.model_config = model_config self.device_config = device_config self.target_device = torch.device(device_config.device) + self.model_streamer_distributed = ( + os.environ.get("MX_MS_DISTRIBUTED", "0").lower() in ("1", "true") + ) def build_identity(self) -> p2p_pb2.SourceIdentity: return build_sglang_source_identity( @@ -187,7 +198,7 @@ def _model_streamer_distributed_enabled(self) -> bool: return ( tp_size > 1 and self.is_cuda_alike() - and os.environ.get("MX_MS_DISTRIBUTED", "0").lower() in ("1", "true") + and self.model_streamer_distributed ) @@ -326,7 +337,9 @@ def build_sglang_load_context( adapter = SglangAdapter(load_config, model_config, device_config) worker_rank = adapter.get_worker_rank() global_rank = adapter.get_global_rank() - server_url = getattr(load_config, "modelexpress_url", None) + server_url = resolve_metadata_server_url( + getattr(load_config, "modelexpress_url", None), + ) return LoadContext( model_config=model_config, load_config=load_config, @@ -340,5 +353,8 @@ def build_sglang_load_context( server_url=server_url, ), worker_id=uuid.uuid4().hex[:8], + metadata_server_url=server_url, + metadata_port=resolve_metadata_port(), + model_streamer_uri=resolve_model_streamer_uri(model_config), adapter=adapter, ) diff --git a/modelexpress_client/python/modelexpress/engines/trtllm/__init__.py b/modelexpress_client/python/modelexpress/engines/trtllm/__init__.py new file mode 100644 index 000000000..fc6a9d6ce --- /dev/null +++ b/modelexpress_client/python/modelexpress/engines/trtllm/__init__.py @@ -0,0 +1,20 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""TensorRT-LLM integration for ModelExpress.""" + +from .adapter import ( + TrtllmAdapter, + TrtllmLoadConfig, + TrtllmModelConfig, + build_trtllm_load_context, +) +from .loader import MXCheckpointLoader + +__all__ = [ + "TrtllmAdapter", + "TrtllmLoadConfig", + "TrtllmModelConfig", + "MXCheckpointLoader", + "build_trtllm_load_context", +] diff --git a/modelexpress_client/python/modelexpress/engines/trtllm/adapter.py b/modelexpress_client/python/modelexpress/engines/trtllm/adapter.py new file mode 100644 index 000000000..3a25f0144 --- /dev/null +++ b/modelexpress_client/python/modelexpress/engines/trtllm/adapter.py @@ -0,0 +1,326 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""TensorRT-LLM implementation of the ModelExpress engine adapter contract.""" + +from __future__ import annotations + +import logging +import uuid +from collections.abc import Callable +from typing import Any + +import torch + +from ... import p2p_pb2 +from ...adapter import EngineAdapter +from ...load_strategy.context import LoadContext, LoadResult +from ...metadata.client_factory import ( + create_metadata_client, + resolve_metadata_port, + resolve_metadata_server_url, +) + +logger = logging.getLogger("modelexpress.engines.trtllm.adapter") + + +class TrtllmAdapter(EngineAdapter): + """Adapter that maps ModelExpress hooks onto TRT-LLM's live model object.""" + + def __init__( + self, + *, + model_name: str, + checkpoint_dir: str | None = None, + model: Any = None, + mapping: Any = None, + native_loader: Callable[[], dict[str, Any]] | None = None, + ): + self.model_name = model_name + self.checkpoint_dir = checkpoint_dir + self.model = model + self.mapping = _resolve_trtllm_mapping(mapping, model) + self.native_loader = native_loader + self.dtype = _resolve_trtllm_dtype(model) + self.quantization = _resolve_trtllm_quantization(model) + self.target_device = torch.device(f"cuda:{self.get_device_id()}") + + def build_identity(self): + return _build_trtllm_identity( + model_name=self.model_name, + tp_size=int(getattr(self.mapping, "tp_size", 1) or 1), + pp_size=int(getattr(self.mapping, "pp_size", 1) or 1), + ep_size=int(getattr(self.mapping, "moe_ep_size", 1) or 1), + dtype=_dtype_to_identity_string(self.dtype), + quantization=self.quantization, + ) + + def get_worker_rank(self) -> int: + return _get_trtllm_worker_rank(self.mapping, self.get_global_rank()) + + def get_global_rank(self) -> int: + if self.mapping is None: + return self.get_device_id() + return int(self.mapping.rank) + + def get_device_id(self) -> int: + try: + return int(torch.cuda.current_device()) + except Exception: + return 0 + + def get_target_device(self) -> torch.device: + return self.target_device + + def is_cuda_alike(self) -> bool: + return torch.cuda.is_available() + + def discover_tensors(self, result: LoadResult) -> dict[str, torch.Tensor]: + model = result.model or self.model + if model is None: + raise RuntimeError("TRT-LLM tensor discovery requires result.model") + tensors, _ = _collect_cuda_param_tensors(model, self.get_device_id()) + return tensors + + def load_via_native(self, result: LoadResult) -> LoadResult: + if self.native_loader is None: + raise RuntimeError("TRT-LLM native fallback loader is not configured") + return LoadResult( + value=self.native_loader(), + model=None, + publishable=False, + metadata=result.metadata, + ) + + def after_rdma_receive(self, result: LoadResult) -> LoadResult: + # TRT-LLM receivers already got weights over P2P and should not be + # advertised as fresh sources before TRT-LLM's post-load lifecycle. + result.publishable = False + return result + + +def build_trtllm_load_context( + *, + model_name: str, + checkpoint_dir: str | None, + model: Any, + mapping: Any = None, + server_url: str | None = None, + native_loader: Callable[[], dict[str, Any]] | None = None, + source_query_timeout_s: int | None = None, +) -> LoadContext: + """Build a LoadContext for TRT-LLM source publication and metadata lookup.""" + server_url = resolve_metadata_server_url(server_url) + adapter = TrtllmAdapter( + model_name=model_name, + checkpoint_dir=checkpoint_dir, + model=model, + mapping=mapping, + native_loader=native_loader, + ) + worker_rank = adapter.get_worker_rank() + return LoadContext( + model_config=TrtllmModelConfig( + model_name, + dtype=adapter.dtype, + quantization=adapter.quantization, + hf_text_config=_resolve_trtllm_hf_text_config(model), + ), + load_config=TrtllmLoadConfig( + source_query_timeout_s=source_query_timeout_s, + ), + target_device=adapter.get_target_device(), + global_rank=adapter.get_global_rank(), + worker_rank=worker_rank, + device_id=adapter.get_device_id(), + identity=adapter.build_identity(), + mx_client=create_metadata_client( + worker_rank=worker_rank, + server_url=server_url, + ), + worker_id=uuid.uuid4().hex[:8], + metadata_server_url=server_url, + metadata_port=resolve_metadata_port(), + adapter=adapter, + ) + + +class TrtllmModelConfig: + def __init__( + self, + model: str, + *, + dtype: Any, + quantization: str, + hf_text_config: Any, + ): + self.model = model + self.model_weights = None + self.hf_text_config = hf_text_config + self.dtype = dtype + self.quantization = quantization + + +class _TrtllmHfTextConfig: + model_type = "unknown" + + +class TrtllmLoadConfig: + def __init__(self, *, source_query_timeout_s: int | None = None): + self.source_query_timeout_s = source_query_timeout_s + + +def _build_trtllm_identity( + model_name: str, + tp_size: int = 1, + pp_size: int = 1, + ep_size: int = 1, + dtype: str = "unknown", + quantization: str = "", +) -> p2p_pb2.SourceIdentity: + from importlib.metadata import version as pkg_version + + try: + mx_version = pkg_version("modelexpress") + except Exception: + mx_version = "0.0.0" + + return p2p_pb2.SourceIdentity( + mx_version=mx_version, + mx_source_type=p2p_pb2.MX_SOURCE_TYPE_WEIGHTS, + model_name=model_name, + backend_framework=p2p_pb2.BACKEND_FRAMEWORK_TRT_LLM, + tensor_parallel_size=tp_size, + pipeline_parallel_size=pp_size, + expert_parallel_size=ep_size, + dtype=dtype, + quantization=quantization, + ) + + +def _resolve_trtllm_dtype(model: Any) -> Any: + """Resolve dtype from TRT-LLM runtime state without inventing a default.""" + for param in _iter_model_parameters(model): + dtype = getattr(param, "dtype", None) + if dtype is not None: + return dtype + + model_config = getattr(model, "model_config", None) + pretrained_config = getattr(model_config, "pretrained_config", None) + runtime_config = getattr(model, "config", None) + for owner, attr in ( + (runtime_config, "torch_dtype"), + (runtime_config, "dtype"), + (pretrained_config, "torch_dtype"), + (pretrained_config, "dtype"), + ): + value = getattr(owner, attr, None) + if value is not None: + return value + return "unknown" + + +def _iter_model_parameters(model: Any): + parameters = getattr(model, "parameters", None) + if not callable(parameters): + return () + try: + return parameters() + except Exception: + return () + + +def _dtype_to_identity_string(dtype: Any) -> str: + if dtype is None: + return "unknown" + return str(dtype).replace("torch.", "") + + +def _resolve_trtllm_quantization(model: Any) -> str: + model_config = getattr(model, "model_config", None) + quantization = getattr(model_config, "quantization", None) + if quantization: + return str(quantization) + + quant_config = getattr(model_config, "quant_config", None) + quant_algo = getattr(quant_config, "quant_algo", None) + if quant_algo: + return str(quant_algo) + + pretrained_config = getattr(model_config, "pretrained_config", None) + hf_quant_config = getattr(pretrained_config, "quantization_config", None) + if isinstance(hf_quant_config, dict): + quant_method = ( + hf_quant_config.get("quant_method") + or hf_quant_config.get("quant_algo") + or hf_quant_config.get("type") + ) + return str(quant_method or "") + if hf_quant_config: + return str(hf_quant_config) + return "" + + +def _resolve_trtllm_hf_text_config(model: Any) -> Any: + model_config = getattr(model, "model_config", None) + pretrained_config = getattr(model_config, "pretrained_config", None) + if pretrained_config is not None: + return pretrained_config + runtime_config = getattr(model, "config", None) + if runtime_config is not None: + return runtime_config + return _TrtllmHfTextConfig() + + +def _resolve_trtllm_mapping(mapping: Any, model: Any) -> Any: + if mapping is not None: + return mapping + model_mapping = getattr(model, "mapping", None) + if model_mapping is not None: + return model_mapping + model_config = getattr(model, "model_config", None) + return getattr(model_config, "mapping", None) + + +def _get_trtllm_worker_rank(mapping: Any, default: int) -> int: + """Return the model-weight shard key for source/target matching.""" + if mapping is None: + return int(default) + + tp_size = int(getattr(mapping, "tp_size", 1) or 1) + rank = int(mapping.rank) + cp_size = int(getattr(mapping, "cp_size", 1) or 1) + tp_cp_size = max(1, tp_size * cp_size) + # TRT-LLM defines pp_rank = rank // (tp_size * cp_size) and + # tp_rank = rank % (tp_size * cp_size) // cp_size. Context-parallel ranks + # split sequence work, not model weights, so source keys only on PP/TP shards. + return (rank // tp_cp_size) * tp_size + (rank % tp_cp_size) // max(1, cp_size) + + +def _collect_cuda_param_tensors( + torch_model: Any, + device_id: int, +) -> tuple[dict[str, Any], int]: + param_tensors = {} + seen_data_ptrs = set() + total_bytes = 0 + for name, param in torch_model.named_parameters(): + if param.device.type != "cuda" or param.device.index != device_id: + continue + tensor = param.data + ptr = tensor.data_ptr() + if ptr in seen_data_ptrs: + logger.debug("Skipping aliased param: %s (ptr=%x)", name, ptr) + continue + seen_data_ptrs.add(ptr) + param_tensors[name] = tensor + total_bytes += tensor.numel() * tensor.element_size() + return param_tensors, total_bytes + + +__all__ = [ + "TrtllmAdapter", + "TrtllmLoadConfig", + "TrtllmModelConfig", + "build_trtllm_load_context", +] diff --git a/modelexpress_client/python/modelexpress/engines/trtllm/loader.py b/modelexpress_client/python/modelexpress/engines/trtllm/loader.py new file mode 100644 index 000000000..3d174d6ba --- /dev/null +++ b/modelexpress_client/python/modelexpress/engines/trtllm/loader.py @@ -0,0 +1,386 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""TensorRT-LLM checkpoint loader backed by ModelExpress strategies.""" + +from __future__ import annotations + +import logging +import os +import sys +import traceback +from contextlib import contextmanager +from pathlib import Path +from typing import Any, Optional, Union + +from ...load_strategy import LoadResult, LoadStrategyChain +from ...metadata.client_factory import resolve_metadata_server_url + +logger = logging.getLogger("modelexpress.engines.trtllm.loader") + + +class _TrtllmStyleFormatter(logging.Formatter): + _LEVEL_TAGS = { + logging.DEBUG: "D", + logging.INFO: "I", + logging.WARNING: "W", + logging.ERROR: "E", + logging.CRITICAL: "C", + } + + def format(self, record: logging.LogRecord) -> str: + record.mx_level_tag = self._LEVEL_TAGS.get(record.levelno, record.levelname[0]) + return super().format(record) + + +@contextmanager +def _rank_log_scope(global_rank: int): + """Mirror TRT-LLM worker logs to stderr and a line-buffered rank log.""" + log_dir = os.environ.get("MX_TRANSFER_LOG_DIR", "/tmp/mx_logs") + os.makedirs(log_dir, exist_ok=True) + rank_log = os.path.abspath(os.path.join(log_dir, f"rank{global_rank}.log")) + + mx_logger = logging.getLogger("modelexpress") + mx_logger.setLevel(logging.INFO) + for handler in list(mx_logger.handlers): + if getattr(handler, "_mx_rank_log_path", None) == rank_log or getattr( + handler, "baseFilename", None + ) == rank_log: + mx_logger.removeHandler(handler) + handler.close() + + formatter = _TrtllmStyleFormatter( + "[%(asctime)s] [ModelExpress] [%(mx_level_tag)s] %(message)s", + datefmt="%m/%d/%Y-%H:%M:%S", + ) + stream = open(rank_log, "a", buffering=1) + file_handler = logging.StreamHandler(stream) + file_handler.setLevel(logging.INFO) + file_handler.setFormatter(formatter) + file_handler._mx_rank_log_path = rank_log + mx_logger.addHandler(file_handler) + + # TRT-LLM MPI workers do not reliably inherit a root handler that forwards + # ModelExpress logs to the container log stream. Keep the durable rank log, + # but also mirror transfer events to stderr so `kubectl logs` shows state. + stderr_handler = logging.StreamHandler(sys.stderr) + stderr_handler.setLevel(logging.INFO) + stderr_handler.setFormatter(formatter) + stderr_handler._mx_trtllm_stderr = True + mx_logger.addHandler(stderr_handler) + try: + yield + finally: + try: + try: + file_handler.flush() + stream.flush() + os.fsync(stream.fileno()) + stderr_handler.flush() + except Exception: + pass + finally: + mx_logger.removeHandler(file_handler) + mx_logger.removeHandler(stderr_handler) + stream.close() + + +def _resolve_mx_model_name( + model_name_arg: Optional[str], + checkpoint_dir: Optional[str], +) -> str: + """Resolve the model identity used for TRT-LLM ModelExpress source matching.""" + if model_name_arg: + return str(model_name_arg) + + env_model_name = os.environ.get("MODEL_NAME") + if env_model_name: + return env_model_name + + if checkpoint_dir: + path = os.path.normpath(str(checkpoint_dir)) + parts = path.split(os.sep) + if ( + len(parts) >= 3 + and parts[-2] == "snapshots" + and parts[-3].startswith("models--") + ): + return parts[-3].removeprefix("models--").replace("--", "/") + return os.path.basename(path) + + return "unknown" + + +try: + from tensorrt_llm._torch.models.checkpoints.base_config_loader import ( + BaseConfigLoader, + ) + from tensorrt_llm._torch.models.checkpoints.base_weight_loader import ( + BaseWeightLoader, + ConsumableWeightsDict, + ) + from tensorrt_llm._torch.models.checkpoints.base_weight_mapper import ( + BaseWeightMapper, + ) + from tensorrt_llm._torch.models.checkpoints.auto_mapper import ( + AutoCheckpointMapper, + ) + from tensorrt_llm._torch.models.checkpoints.hf.checkpoint_loader import ( + HfCheckpointLoader, + ) + from tensorrt_llm.mapping import Mapping +except Exception as exc: + _TRTLLM_IMPORT_ERROR = exc + + class MXCheckpointLoader: + """Placeholder used when TensorRT-LLM is not installed.""" + + def __init__(self, *args, **kwargs): + raise ImportError( + "modelexpress.engines.trtllm.loader.MXCheckpointLoader " + "requires TensorRT-LLM to be installed." + ) from _TRTLLM_IMPORT_ERROR + +else: + + class MXCheckpointLoader(HfCheckpointLoader): + """TRT-LLM checkpoint loader backed by ModelExpress load strategies.""" + + def __init__( + self, + *, + weight_loader: Optional[BaseWeightLoader] = None, + weight_mapper: Optional[BaseWeightMapper] = None, + config_loader: Optional[BaseConfigLoader] = None, + mx_server_url: Optional[str] = None, + model_name: Optional[Union[str, Path]] = None, + query_timeout_s: Optional[int] = None, + ): + super().__init__( + weight_loader=weight_loader, + weight_mapper=weight_mapper, + config_loader=config_loader, + ) + self._checkpoint_format = "modelexpress" + self._mx_server_url = mx_server_url + 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._last_load_ctx = None + + @property + def checkpoint_format(self) -> str: + return "modelexpress" + + @property + def mx_server_url(self) -> Optional[str]: + return self._mx_server_url + + @property + def model_name(self) -> Optional[str]: + return self._model_name + + @property + def query_timeout_s(self) -> Optional[int]: + return self._query_timeout_s + + @property + def p2p_succeeded(self) -> bool: + return self._p2p_succeeded + + def is_weights_preloaded(self) -> bool: + return self._p2p_succeeded + + def load_weights( + self, + checkpoint_dir: str, + mapping: Mapping, + *, + model=None, + **kwargs, + ) -> dict[str, Any]: + """Load TRT-LLM weights through ModelExpress, falling back to native disk. + + `mapping` matches TRT-LLM's BaseCheckpointLoader contract. `model` + is a ModelExpress-only extension passed by TRT-LLM's model loader + so P2P can write directly into live parameter buffers. Keep `model` + explicit instead of inside `kwargs` so native fallback loaders + never receive an argument they do not understand. + """ + fallback_loader = lambda: self._load_from_disk( + checkpoint_dir, + mapping, + **kwargs, + ) + + if model is None: + logger.info( + "TRT-LLM ModelExpress loader did not receive a model reference; " + "using native checkpoint loading." + ) + return fallback_loader() + + self._p2p_succeeded = False + self._last_load_ctx = None + try: + ctx = self._build_load_context( + model=model, + checkpoint_dir=checkpoint_dir, + mapping=mapping, + native_loader=fallback_loader, + source_query_timeout_s=self._query_timeout_s, + ) + self._last_load_ctx = ctx + with _rank_log_scope(ctx.global_rank): + result = LoadStrategyChain.run(model, ctx) + except Exception: + logger.warning( + "ModelExpress strategy chain failed; falling back to native checkpoint " + "loading.\n%s", + traceback.format_exc(), + ) + return fallback_loader() + + selected_strategy = getattr(ctx, "selected_strategy", None) + if selected_strategy != "rdma": + if result is None: + result = {} + if not isinstance(result, (dict, ConsumableWeightsDict)): + raise TypeError( + "TRT-LLM native fallback must return a weight dict " + "or ConsumableWeightsDict " + f"or None, got {type(result).__name__}" + ) + if result: + fallback_bytes = sum( + tensor.numel() * tensor.element_size() + for tensor in result.values() + ) + logger.info( + "ModelExpress default strategy loaded %d fallback weights " + "(%.2f MiB) through TRT-LLM native checkpoint loading.", + len(result), + fallback_bytes / (1 << 20), + ) + return result + + self._p2p_succeeded = True + logger.info( + "ModelExpress P2P weight transfer succeeded from %s", + self._mx_server_url, + ) + return {} + + def get_initialized_weight_mapper(self, model, config): + """Use TRT-LLM's HF mapper for ModelExpress disk fallback weights. + + ModelExpress is a transport format here. When the strategy chain + falls back to disk, the checkpoint tensors are still HF-formatted, + so asking TRT-LLM for a transport-specific mapper can miss + architecture-specific mappers in released TRT-LLM images. + """ + if self.weight_mapper is not None: + self.weight_mapper.init_model_and_config(model, config) + return self.weight_mapper + + if config.pretrained_config and config.pretrained_config.architectures: + model_arch = config.pretrained_config.architectures[0] + else: + raise ValueError("Cannot determine model architecture from config") + weight_mapper = AutoCheckpointMapper.get("HF", model_arch) + weight_mapper.init_model_and_config(model, config) + self.weight_mapper = weight_mapper + return weight_mapper + + def _load_from_disk( + self, + checkpoint_dir: str, + mapping: Mapping, + **kwargs, + ) -> dict[str, Any]: + return super().load_weights(checkpoint_dir, mapping=mapping, **kwargs) + + def publish_as_source( + self, + model, + checkpoint_dir: Optional[str] = None, + ) -> None: + self._publish_as_source(model, checkpoint_dir=checkpoint_dir) + + def _build_load_context( + self, + *, + model, + checkpoint_dir: Optional[str], + mapping: Optional[Mapping] = None, + native_loader=None, + source_query_timeout_s: Optional[int] = None, + ): + from .adapter import build_trtllm_load_context + + server_url = resolve_metadata_server_url(self._mx_server_url) + resolved_name = _resolve_mx_model_name(self._model_name, checkpoint_dir) + return build_trtllm_load_context( + model_name=resolved_name, + checkpoint_dir=checkpoint_dir, + model=model, + mapping=mapping, + server_url=server_url, + native_loader=native_loader, + source_query_timeout_s=source_query_timeout_s, + ) + + def _publish_as_source( + self, + model, + *, + checkpoint_dir: Optional[str] = None, + mapping: Optional[Mapping] = None, + ) -> None: + self._publish_current_model(model) + + def _publish_current_model(self, model) -> None: + ctx = self._last_load_ctx + if ctx is None or getattr(ctx.adapter, "model", None) is not model: + logger.warning( + "Skipping ModelExpress source publish because TRT-LLM did not provide " + "a matching load context for this model." + ) + return + try: + from ...load_strategy import publish_loaded_model + + with _rank_log_scope(ctx.global_rank): + publish_loaded_model(LoadResult(value=model, model=model), ctx) + except Exception: + logger.warning( + "Failed to publish weights to ModelExpress server at %s.\n%s", + self._mx_server_url, + traceback.format_exc(), + ) + + def post_load_publish( + self, + model, + *, + checkpoint_dir: str, + weights_preloaded: bool = False, + ) -> None: + """Publish only after TRT-LLM has finished its post-load path. + + load_weights() may return native fallback weights for TRT-LLM to + merge through its own model.load_weights() path. Publishing inside + the strategy chain would advertise stale tensors before TRT-LLM + post-processing. RDMA receivers also publish here: after P2P writes + weights into live buffers and TRT-LLM completes post-load setup, + the receiver can safely become another source. + """ + # TRT-LLM applies fallback weights after load_weights() returns. + # Reusing the load context keeps disk fallback and RDMA receiver + # publication under the same worker_id used by the load attempt. + self._publish_current_model(model) + + +__all__ = [ + "MXCheckpointLoader", +] diff --git a/modelexpress_client/python/modelexpress/engines/vllm/adapter.py b/modelexpress_client/python/modelexpress/engines/vllm/adapter.py index f3b0ae53e..0e0aed616 100644 --- a/modelexpress_client/python/modelexpress/engines/vllm/adapter.py +++ b/modelexpress_client/python/modelexpress/engines/vllm/adapter.py @@ -14,8 +14,16 @@ import torch from ...adapter import EngineAdapter -from ...load_strategy.context import LoadContext, LoadResult -from ...metadata.client_factory import create_metadata_client +from ...load_strategy.context import ( + LoadContext, + LoadResult, + resolve_model_streamer_uri, +) +from ...metadata.client_factory import ( + create_metadata_client, + resolve_metadata_port, + resolve_metadata_server_url, +) from ...metadata.publish import build_source_identity from ...rank_utils import get_global_rank from ...tensor_utils import adopt_hidden_tensors, capture_tensor_attrs, collect_module_tensors @@ -34,6 +42,9 @@ def __init__(self, vllm_config, model_config): self.model_config = model_config self.load_config = vllm_config.load_config self.target_device = self._resolve_target_device() + self.model_streamer_distributed = ( + os.environ.get("MX_MS_DISTRIBUTED", "0").lower() in ("1", "true") + ) def build_identity(self): return build_source_identity(self.vllm_config, self.model_config) @@ -188,7 +199,7 @@ def _model_streamer_distributed_enabled(self) -> bool: tp_size = getattr(self.vllm_config.parallel_config, "tensor_parallel_size", 1) return ( tp_size > 1 - and os.environ.get("MX_MS_DISTRIBUTED", "0").lower() in ("1", "true") + and self.model_streamer_distributed ) @@ -240,6 +251,7 @@ def build_vllm_load_context(vllm_config, model_config) -> LoadContext: adapter = VllmAdapter(vllm_config, model_config) global_rank = adapter.get_global_rank() worker_rank = adapter.get_worker_rank() + server_url = resolve_metadata_server_url() return LoadContext( model_config=model_config, load_config=vllm_config.load_config, @@ -248,7 +260,10 @@ def build_vllm_load_context(vllm_config, model_config) -> LoadContext: worker_rank=worker_rank, device_id=adapter.get_device_id(), identity=adapter.build_identity(), - mx_client=create_metadata_client(worker_rank=worker_rank), + mx_client=create_metadata_client(worker_rank=worker_rank, server_url=server_url), worker_id=uuid.uuid4().hex[:8], + metadata_server_url=server_url, + metadata_port=resolve_metadata_port(), + model_streamer_uri=resolve_model_streamer_uri(model_config), adapter=adapter, ) diff --git a/modelexpress_client/python/modelexpress/load_strategy/__init__.py b/modelexpress_client/python/modelexpress/load_strategy/__init__.py index dfeedcb49..c644a0149 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/__init__.py +++ b/modelexpress_client/python/modelexpress/load_strategy/__init__.py @@ -25,6 +25,7 @@ publish_source_if_supported, register_tensors, publish_metadata, + publish_loaded_model, unpublish_metadata, ) @@ -36,6 +37,7 @@ "SourceTransferError", "register_tensors", "publish_metadata", + "publish_loaded_model", "unpublish_metadata", ] @@ -85,6 +87,7 @@ def run(model: nn.Module, ctx: LoadContext) -> nn.Module: logger.info(f"[Worker {ctx.global_rank}] Trying strategy: {strategy.name}") try: result = strategy.load(result, ctx) + ctx.selected_strategy = strategy.name publish_source_if_supported(result, ctx) span.set_attribute("weight_loading_strategy", strategy.name) return result.value diff --git a/modelexpress_client/python/modelexpress/load_strategy/base.py b/modelexpress_client/python/modelexpress/load_strategy/base.py index 85306fe87..a69a6d700 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/base.py +++ b/modelexpress_client/python/modelexpress/load_strategy/base.py @@ -6,12 +6,10 @@ from __future__ import annotations import logging -import os import uuid from abc import ABC, abstractmethod from typing import TYPE_CHECKING, ClassVar -import torch import torch.nn as nn from ..nixl_transfer import is_nixl_available @@ -107,8 +105,7 @@ def _as_load_result(result_or_model: LoadResult | nn.Module) -> LoadResult: def _metadata_publication_configured(ctx: LoadContext) -> bool: """Return whether this worker has a metadata path for P2P serving.""" - server_addr = os.environ.get("MODEL_EXPRESS_URL") or os.environ.get("MX_SERVER_ADDRESS") - if server_addr: + if getattr(ctx, "metadata_server_url", None): return True return getattr(ctx.mx_client, "REQUIRES_P2P_METADATA", False) is True @@ -143,8 +140,7 @@ def register_tensors(result_or_model: LoadResult | nn.Module, ctx: LoadContext) log_tensor_summary(ctx.tensors, ctx.global_rank, "Registering tensors") if ctx.nixl_manager is None: - base_port = int(os.environ.get("MX_METADATA_PORT", "5555")) - listen_port = base_port + ctx.device_id + listen_port = ctx.metadata_port + ctx.device_id ctx.nixl_manager = _init_nixl_manager( ctx.global_rank, ctx.device_id, "auto", listen_port ) @@ -174,11 +170,11 @@ def publish_metadata(ctx: LoadContext) -> None: f"[Worker {ctx.global_rank}] No NIXL manager, skipping metadata publish" ) return - # Decentralized backends (k8s-service) have no central server - # address; their metadata path is entirely peer-to-peer. - # Only bail on missing MODEL_EXPRESS_URL / MX_SERVER_ADDRESS when the - # client actually needs a central coordinator. Strict `is True` - # check so MagicMock's auto-attribute doesn't masquerade as the flag. + # Decentralized backends (k8s-service) have no central server address; + # their metadata path is entirely peer-to-peer. Only bail on missing + # metadata_server_url when the client actually needs a central coordinator. + # Strict `is True` check so MagicMock's auto-attribute doesn't masquerade as + # the flag. if not _metadata_publication_configured(ctx): logger.info( f"[Worker {ctx.global_rank}] No MX server configured, skipping metadata publish" @@ -196,11 +192,46 @@ def publish_metadata(ctx: LoadContext) -> None: ) +def publish_loaded_model(result: LoadResult, ctx: LoadContext) -> None: + """Register and publish an already-loaded model as a reusable source. + + Some engine lifecycles load weights outside the strategy chain and only + have a publish callback after the model is ready. Keep the source context + reachable from the model so its NIXL manager, identity, and metadata client + remain alive while the engine owns the model object. + """ + if not result.publishable: + return + + register_tensors(result, ctx) + publish_metadata(ctx) + + model = result.model + if model is None: + return + + _retain_source_runtime(model, ctx) + + +def _retain_source_runtime(model: nn.Module, ctx: LoadContext) -> None: + """Keep source-serving runtime objects alive for the model lifetime. + + Publishing metadata advertises this process as a live source, but the actual + source still depends on local runtime state. The context owns the NIXL + manager, registered memory endpoints, worker identity, and metadata client. + Some engine lifecycles, notably TRT-LLM post-load publish, create this + context inside a short callback rather than storing it on a ModelExpress + loader instance. Attach the current context to the model so Python does not + collect the source runtime while the engine is still serving the model. + """ + setattr(model, "_mx_load_context", ctx) + + def publish_source_if_supported(result: LoadResult, ctx: LoadContext) -> None: """Best-effort source publication after a successful load.""" if result.model_for_publish is None: return - publish_metadata(ctx) + publish_loaded_model(result, ctx) def unpublish_metadata(ctx: LoadContext) -> None: diff --git a/modelexpress_client/python/modelexpress/load_strategy/context.py b/modelexpress_client/python/modelexpress/load_strategy/context.py index fab53c567..b294e95fa 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/context.py +++ b/modelexpress_client/python/modelexpress/load_strategy/context.py @@ -5,6 +5,7 @@ from __future__ import annotations +import os from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Generic, TypeAlias, TypeVar @@ -16,6 +17,10 @@ if TYPE_CHECKING: from ..adapter import EngineAdapter + from ..engines.trtllm.adapter import ( + TrtllmLoadConfig, + TrtllmModelConfig, + ) from ..nixl_transfer import NixlTransferManager from ..vmm import VmmArena from sglang.srt.configs.load_config import LoadConfig as SglangLoadConfig @@ -23,8 +28,12 @@ from vllm.config import ModelConfig as VllmModelConfig from vllm.config.load import LoadConfig as VllmLoadConfig - EngineModelConfig: TypeAlias = VllmModelConfig | SglangModelConfig - EngineLoadConfig: TypeAlias = VllmLoadConfig | SglangLoadConfig + # vLLM and SGLang expose native config objects. TRT-LLM does not expose the + # same load-strategy shape here, so its adapter provides a small shim. + EngineModelConfig: TypeAlias = ( + VllmModelConfig | SglangModelConfig | TrtllmModelConfig + ) + EngineLoadConfig: TypeAlias = VllmLoadConfig | SglangLoadConfig | TrtllmLoadConfig else: EngineModelConfig = object EngineLoadConfig = object @@ -60,7 +69,11 @@ class LoadContext: identity: p2p_pb2.SourceIdentity mx_client: MxClientBase worker_id: str + metadata_server_url: str | None = None + metadata_port: int = 5555 + model_streamer_uri: str | None = None adapter: EngineAdapter | None = None + selected_strategy: str | None = None nixl_manager: NixlTransferManager | None = None tensors: dict[str, torch.Tensor] = field(default_factory=dict) # When MX_VMM_ARENA=1, maybe_enter_vmm_arena populates this with the @@ -69,3 +82,16 @@ class LoadContext: # cuMemGetHandleForAddressRange + ibv_reg_dmabuf_mr, collapsing # O(plugin_calls) MRs to 1. vmm_arena: VmmArena | None = None + + +def resolve_model_streamer_uri(model_config: EngineModelConfig) -> str | None: + """Resolve ModelStreamer URI while preserving MX_MODEL_URI gate semantics.""" + if not os.environ.get("MX_MODEL_URI"): + return None + model_weights = getattr(model_config, "model_weights", None) + if model_weights: + return str(model_weights) + model = getattr(model_config, "model", None) + if model: + return str(model) + return None diff --git a/modelexpress_client/python/modelexpress/load_strategy/default_strategy.py b/modelexpress_client/python/modelexpress/load_strategy/default_strategy.py index ef9b30974..48aaf56f9 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/default_strategy.py +++ b/modelexpress_client/python/modelexpress/load_strategy/default_strategy.py @@ -34,5 +34,6 @@ def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: except Exception as e: raise StrategyFailed(str(e), mutated=True) from e - register_tensors(result, ctx) + if result.model_for_publish is not None: + register_tensors(result, ctx) return result diff --git a/modelexpress_client/python/modelexpress/load_strategy/model_streamer_strategy.py b/modelexpress_client/python/modelexpress/load_strategy/model_streamer_strategy.py index 49968ed4c..09f4e07cd 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/model_streamer_strategy.py +++ b/modelexpress_client/python/modelexpress/load_strategy/model_streamer_strategy.py @@ -12,7 +12,6 @@ import importlib.util import logging -import os from typing import Iterator import torch @@ -27,11 +26,9 @@ class ModelStreamerStrategy(LoadStrategy): """Load weights by streaming safetensors via runai-model-streamer. - Activated by setting MX_MODEL_URI (gate only). The actual URI used to - stream weights is taken from `model_config.model_weights` if set - (object-storage URIs: s3://, gs://, az://) and falls back to - `model_config.model` otherwise (local paths, HF-resolved snapshots). - Engine adapters provide the concrete ModelStreamer iterator. + Activated by LoadContext.model_streamer_uri. Context builders preserve the + legacy MX_MODEL_URI gate behavior and freeze the resolved model path before + the strategy chain runs. """ name = "model_streamer" @@ -49,17 +46,18 @@ def is_available(self, ctx: LoadContext) -> bool: ) return False - model_uri = os.environ.get("MX_MODEL_URI", "") - if not model_uri: + if not ctx.model_streamer_uri: logger.info( - f"[Worker {ctx.global_rank}] MX_MODEL_URI not set, skipping model streamer" + f"[Worker {ctx.global_rank}] ModelStreamer URI not configured, skipping" ) return False return True def load(self, result: LoadResult, ctx: LoadContext) -> LoadResult: result = _as_load_result(result) - model_uri = getattr(ctx.model_config, "model_weights", None) or ctx.model_config.model + model_uri = ctx.model_streamer_uri + if not model_uri: + raise StrategyFailed("ModelStreamer URI not configured", mutated=False) logger.info(f"[Worker {ctx.global_rank}] Attempting model streamer loading from {model_uri}") try: diff --git a/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py b/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py index 086235a14..41bdca2e7 100644 --- a/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py +++ b/modelexpress_client/python/modelexpress/load_strategy/rdma_strategy.py @@ -6,7 +6,6 @@ from __future__ import annotations import logging -import os import random import time @@ -29,6 +28,16 @@ logger = logging.getLogger("modelexpress.strategy_rdma") MAX_SOURCE_RETRIES = 3 +SOURCE_QUERY_POLL_INTERVAL_S = 5.0 + + +def _source_query_timeout_s(ctx: LoadContext) -> float | None: + timeout_s = getattr(ctx.load_config, "source_query_timeout_s", None) + if timeout_s is None: + return None + if not isinstance(timeout_s, (int, float)): + return None + return float(timeout_s) class RdmaStrategy(LoadStrategy): @@ -63,9 +72,8 @@ def is_available(self, ctx: LoadContext) -> bool: # metadata; skip the central-server precondition for them. # Strict `is True` check so MagicMock's auto-attribute doesn't # masquerade as the flag in tests. - server_addr = os.environ.get("MODEL_EXPRESS_URL") or os.environ.get("MX_SERVER_ADDRESS") requires_p2p = getattr(ctx.mx_client, "REQUIRES_P2P_METADATA", False) is True - if not server_addr and not requires_p2p: + if not ctx.metadata_server_url and not requires_p2p: logger.info(f"[Worker {ctx.global_rank}] No MX server configured, skipping RDMA") return False @@ -137,6 +145,27 @@ def _find_source_instances( self, ctx: LoadContext, ) -> list[p2p_pb2.SourceInstanceRef]: """Return all READY source instances (shuffled for load balancing).""" + timeout_s = _source_query_timeout_s(ctx) + deadline = None if timeout_s is None else time.monotonic() + timeout_s + + while True: + candidates = self._list_source_instances(ctx) + if candidates or deadline is None: + return candidates + + remaining_s = deadline - time.monotonic() + if remaining_s <= 0: + logger.info( + f"[Worker {ctx.global_rank}] No RDMA source found within " + f"{timeout_s}s, falling through" + ) + return [] + + time.sleep(min(SOURCE_QUERY_POLL_INTERVAL_S, remaining_s)) + + def _list_source_instances( + self, ctx: LoadContext, + ) -> list[p2p_pb2.SourceInstanceRef]: try: list_resp = ctx.mx_client.list_sources( identity=ctx.identity, diff --git a/modelexpress_client/python/modelexpress/metadata/client_factory.py b/modelexpress_client/python/modelexpress/metadata/client_factory.py index 58f5b2927..1056218a8 100644 --- a/modelexpress_client/python/modelexpress/metadata/client_factory.py +++ b/modelexpress_client/python/modelexpress/metadata/client_factory.py @@ -23,7 +23,7 @@ import logging import os -from ..client import MxClient, MxClientBase +from ..client import MxClient, MxClientBase, _parse_server_address from .k8s_service_client import MxK8sServiceClient logger = logging.getLogger("modelexpress.metadata.client_factory") @@ -34,6 +34,25 @@ _K8S_SERVICE_ALIASES = frozenset({"k8s-service", "service"}) +def resolve_metadata_server_url(server_url: str | None = None) -> str | None: + """Return the explicitly configured central metadata server, if any. + + Unlike MxClient's connection resolver, this intentionally does not fall + back to localhost. Strategies use this value as a configuration signal for + whether central metadata was requested at all. + """ + if server_url: + return _parse_server_address(server_url) + url = os.environ.get("MODEL_EXPRESS_URL") or os.environ.get("MX_SERVER_ADDRESS") + if not url: + return None + return _parse_server_address(url) + + +def resolve_metadata_port() -> int: + return int(os.environ.get("MX_METADATA_PORT", "5555")) + + def create_metadata_client( worker_rank: int | None = None, server_url: str | None = None, diff --git a/modelexpress_client/python/modelexpress/trtllm_live_transfer.py b/modelexpress_client/python/modelexpress/trtllm_live_transfer.py deleted file mode 100644 index e84d3f0b9..000000000 --- a/modelexpress_client/python/modelexpress/trtllm_live_transfer.py +++ /dev/null @@ -1,626 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -""" -ModelExpress Live Model P2P Transfer for TensorRT-LLM. - -Transfers model weights directly between running TRT-LLM instances via NIXL RDMA. -Source registers its model parameter GPU buffers; target receives into its own -model parameter buffers. No format conversion, no disk I/O, no CPU round-trip. - -Target usage (via checkpoint_loader): - from modelexpress.trtllm_live_transfer import MxLiveCheckpointLoader - loader = MxLiveCheckpointLoader() - llm = LLM(model="Llama-70B", checkpoint_loader=loader, - load_format=LoadFormat.PRESHARDED, tp=8) -""" - -from __future__ import annotations - -import logging -import os -import time -import uuid -from typing import Any, Optional - -import torch - -from .client import MxClient -from . import p2p_pb2 - -logger = logging.getLogger("modelexpress.trtllm_live_transfer") - - -def _build_trtllm_identity( - model_name: str, - tp_size: int = 1, - ep_size: int = 1, - dtype: str = "bfloat16", -) -> p2p_pb2.SourceIdentity: - from importlib.metadata import version as pkg_version - - try: - mx_version = pkg_version("modelexpress") - except Exception: - mx_version = "0.0.0" - - return p2p_pb2.SourceIdentity( - mx_version=mx_version, - mx_source_type=p2p_pb2.MX_SOURCE_TYPE_WEIGHTS, - model_name=model_name, - backend_framework=p2p_pb2.BACKEND_FRAMEWORK_TRT_LLM, - tensor_parallel_size=tp_size, - pipeline_parallel_size=1, - expert_parallel_size=ep_size, - dtype=dtype, - ) - - - -def publish_model_params(torch_model: Any) -> None: - """Publish this rank's model params to ModelExpress directly from a torch model. - - Called from ModelLoader.load() BEFORE post_load_weights() so that targets - receive pre-processed weights and can run their own post_load_weights(). - - Each rank publishes independently via MxClient (per-worker API). - """ - from .nixl_transfer import NixlTransferManager - - if not hasattr(torch_model, "named_parameters"): - logger.warning("publish_model_params: model has no named_parameters") - return - - device_id = torch.cuda.current_device() - try: - from mpi4py import MPI - mpi_rank = MPI.COMM_WORLD.Get_rank() - except Exception: - mpi_rank = device_id - - model_name = os.environ.get("MODEL_NAME", "unknown") - mx_server = os.environ.get("MODEL_EXPRESS_URL", "modelexpress-server:8001") - - param_tensors = {} - seen_data_ptrs = set() - total_bytes = 0 - for name, param in torch_model.named_parameters(): - if param.device.type == "cuda" and param.device.index == device_id: - ptr = param.data.data_ptr() - if ptr in seen_data_ptrs: - logger.debug("Skipping aliased param: %s (ptr=%x)", name, ptr) - continue - seen_data_ptrs.add(ptr) - param_tensors[name] = param.data - total_bytes += param.numel() * param.element_size() - - if not param_tensors: - logger.warning("publish_model_params: no params on device %d (rank %d)", device_id, mpi_rank) - return - - logger.info( - "ModelExpress publish_model_params: '%s' rank %d (GPU %d), %d params, %.2f GB (PRE post_load_weights)", - model_name, mpi_rank, device_id, len(param_tensors), total_bytes / 1e9, - ) - - nixl_mgr = NixlTransferManager( - agent_name=f"trtllm-live-source-rank{mpi_rank}-{os.getpid()}", - device_id=device_id, - ) - nixl_mgr.initialize() - nixl_mgr.register_tensors(param_tensors) - - if not hasattr(torch_model, '_mx_nixl_managers'): - torch_model._mx_nixl_managers = [] - torch_model._mx_nixl_managers.append(nixl_mgr) - - tensor_protos = [ - p2p_pb2.TensorDescriptor( - name=name, - addr=tensor.data_ptr(), - size=tensor.numel() * tensor.element_size(), - device_id=device_id, - dtype=str(tensor.dtype), - ) - for name, tensor in param_tensors.items() - ] - - worker = p2p_pb2.WorkerMetadata( - worker_rank=mpi_rank, - nixl_metadata=nixl_mgr.nixl_metadata, - tensors=tensor_protos, - ) - - identity = _build_trtllm_identity(model_name=model_name) - worker_id = uuid.uuid4().hex[:8] - mx_client = MxClient(server_url=mx_server) - try: - mx_source_id = mx_client.publish_metadata( - identity=identity, worker=worker, worker_id=worker_id, - ) - - logger.info( - "ModelExpress worker rank %d (GPU %d) published %.2f GB (mx_source_id=%s)", - mpi_rank, device_id, total_bytes / 1e9, mx_source_id, - ) - finally: - mx_client.close() - - -def publish_from_worker(worker: Any) -> None: - """Publish this rank's model params to ModelExpress from inside a TRT-LLM executor worker. - - Call this from TensorRT-LLM's BaseWorker.setup_engine() after the engine is created, - when MODEL_EXPRESS_SOURCE=1. The worker process has the real model (worker.engine.model_engine.model). - Each rank publishes its own NIXL metadata and tensor descriptors to the MX server. - - Requires patching TRT-LLM's base_worker.setup_engine to call this at the end, e.g.: - - if os.environ.get("MODEL_EXPRESS_SOURCE"): - try: - from modelexpress.trtllm_live_transfer import publish_from_worker - publish_from_worker(self) - except Exception as e: - logger.warning("ModelExpress publish_from_worker failed: %s", e) - """ - from .nixl_transfer import NixlTransferManager - - engine = getattr(worker, "engine", None) - if engine is None: - logger.warning("publish_from_worker: worker has no engine") - return - model_engine = getattr(engine, "model_engine", None) - if model_engine is None: - logger.warning("publish_from_worker: engine has no model_engine (not PyExecutor?)") - return - torch_model = getattr(model_engine, "model", None) - if torch_model is None or not hasattr(torch_model, "named_parameters"): - logger.warning("publish_from_worker: model_engine has no torch model") - return - - device_id = torch.cuda.current_device() - try: - from mpi4py import MPI - mpi_rank = MPI.COMM_WORLD.Get_rank() - except Exception: - mpi_rank = getattr(worker, "rank", device_id) - - model_name = os.environ.get("MODEL_NAME", "unknown") - mx_server = os.environ.get("MODEL_EXPRESS_URL", "modelexpress-server:8001") - - param_tensors = {} - seen_data_ptrs = set() - total_bytes = 0 - for name, param in torch_model.named_parameters(): - if param.device.type == "cuda" and param.device.index == device_id: - ptr = param.data.data_ptr() - if ptr in seen_data_ptrs: - logger.debug("Skipping aliased param: %s (ptr=%x)", name, ptr) - continue - seen_data_ptrs.add(ptr) - param_tensors[name] = param.data - total_bytes += param.numel() * param.element_size() - - if not param_tensors: - logger.warning("publish_from_worker: no params on device %d (rank %d)", device_id, mpi_rank) - return - - logger.info( - "ModelExpress worker publish: '%s' rank %d (GPU %d), %d params, %.2f GB", - model_name, mpi_rank, device_id, len(param_tensors), total_bytes / 1e9, - ) - - if logger.isEnabledFor(logging.DEBUG): - for name, tensor in list(param_tensors.items())[:5]: - val = tensor.to(torch.float32) - cksum = val.sum().item() - nonzero = (tensor != 0).sum().item() - logger.debug( - "SOURCE CHECKSUM rank %d: %s shape=%s dtype=%s sum=%.4f nonzero=%d/%d", - mpi_rank, name, list(tensor.shape), tensor.dtype, - cksum, nonzero, tensor.numel(), - ) - - nixl_mgr = NixlTransferManager( - agent_name=f"trtllm-live-source-rank{mpi_rank}-{os.getpid()}", - device_id=device_id, - ) - nixl_mgr.initialize() - nixl_mgr.register_tensors(param_tensors) - - worker._mx_nixl_manager = nixl_mgr - - tensor_protos = [ - p2p_pb2.TensorDescriptor( - name=name, - addr=tensor.data_ptr(), - size=tensor.numel() * tensor.element_size(), - device_id=device_id, - dtype=str(tensor.dtype), - ) - for name, tensor in param_tensors.items() - ] - - my_worker = p2p_pb2.WorkerMetadata( - worker_rank=mpi_rank, - nixl_metadata=nixl_mgr.nixl_metadata, - tensors=tensor_protos, - ) - - identity = _build_trtllm_identity(model_name=model_name) - worker_id = uuid.uuid4().hex[:8] - mx_client = MxClient(server_url=mx_server) - mx_source_id = mx_client.publish_metadata( - identity=identity, worker=my_worker, worker_id=worker_id, - ) - mx_client.close() - - logger.info( - "ModelExpress worker rank %d (GPU %d) published %.2f GB (mx_source_id=%s)", - mpi_rank, device_id, total_bytes / 1e9, mx_source_id, - ) - - -class MxLiveWeightLoader: - """ - Loads weights via NIXL RDMA directly into model parameter buffers. - - When source publishes TRT-LLM-format param names (from a live model), - this loader matches target params by name and does direct GPU→GPU RDMA. - No format conversion, no fusing, no CPU round-trip. - """ - - def __init__(self, mx_server: Optional[str] = None): - self._source_meta = None - self._mx_server = mx_server - - def load_weights( - self, - checkpoint_dir: str, - mapping: Any = None, - model: Any = None, - **kwargs, - ) -> dict[str, Any]: - from .nixl_transfer import NixlTransferManager - from .types import TensorDescriptor - - # Use provided URL, then env var, then default - mx_server = self._mx_server or os.environ.get("MODEL_EXPRESS_URL") or os.environ.get("MX_SERVER_ADDRESS", "localhost:8001") - model_name = os.environ.get("MODEL_NAME", os.path.basename(checkpoint_dir)) - - if model is None: - raise RuntimeError( - "MxLiveWeightLoader requires model reference. " - "Use load_format=LoadFormat.PRESHARDED to pass model." - ) - - device_id = torch.cuda.current_device() - - # MPI rank may differ from local GPU index in multinode (e.g., rank 4 - # on node B sees local GPU 0). Use MPI rank for source worker matching, - # local GPU index for NIXL agent and tensor registration. - try: - from mpi4py import MPI - mpi_rank = MPI.COMM_WORLD.Get_rank() - except Exception: - mpi_rank = device_id - - # MPI workers' stdout is swallowed by TRT-LLM — write to per-rank file - log_dir = os.environ.get("MX_TRANSFER_LOG_DIR", "/tmp/mx_logs") - os.makedirs(log_dir, exist_ok=True) - rank_log = os.path.join(log_dir, f"rank{mpi_rank}.log") - fh = logging.FileHandler(rank_log, mode="w") - fh.setLevel(logging.INFO) - fh.setFormatter(logging.Formatter("%(asctime)s %(name)s %(levelname)s %(message)s")) - logging.getLogger("modelexpress").addHandler(fh) - - logger.info( - "Live transfer: loading '%s' rank %d (GPU %d)", model_name, mpi_rank, device_id - ) - - # 1. Query source metadata - query_timeout = int(os.environ.get("MX_SOURCE_QUERY_TIMEOUT", "3600")) - source_meta = self._query_source(mx_server, model_name, timeout=query_timeout) - - # Find my rank's source worker — use MPI rank, not local GPU index - my_workers = [w for w in source_meta.workers if w.worker_rank == mpi_rank] - if not my_workers: - raise RuntimeError( - f"No source worker for rank {mpi_rank} (device_id={device_id}). " - f"Source has workers: {[w.worker_rank for w in source_meta.workers]}" - ) - source_worker = my_workers[0] - - # 2. Build name→param map from target model - target_params = {} - for name, param in model.named_parameters(): - if param.device.index == device_id: - target_params[name] = param.data - - logger.info( - "Target has %d params on GPU %d", len(target_params), device_id - ) - - # 3. Build source name→descriptor map - source_descs = {t.name: t for t in source_worker.tensors} - - # 4. Match source and target by name - matched = [] - dtype_cast_needed = [] - unmatched_source = [] - for src_name, src_desc in source_descs.items(): - if src_name in target_params: - dst_param = target_params[src_name] - src_size = src_desc.size - dst_size = dst_param.numel() * dst_param.element_size() - if src_size == dst_size: - matched.append((src_name, src_desc, dst_param)) - else: - # Check if element count matches but dtype differs - src_dtype_str = src_desc.dtype - src_elem_size = 2 if "bfloat16" in src_dtype_str or "float16" in src_dtype_str else 4 if "float32" in src_dtype_str else 1 - src_numel = src_size // src_elem_size if src_elem_size > 0 else 0 - if src_numel == dst_param.numel(): - logger.info( - "Dtype mismatch for %s: source=%s(%d bytes) target=%s(%d bytes) — will cast after transfer", - src_name, src_dtype_str, src_size, dst_param.dtype, dst_size, - ) - dtype_cast_needed.append((src_name, src_desc, dst_param, src_dtype_str)) - else: - logger.warning( - "Size mismatch for %s: source=%d target=%d (numel src=%d dst=%d)", - src_name, src_size, dst_size, src_numel, dst_param.numel(), - ) - else: - unmatched_source.append(src_name) - - if unmatched_source: - logger.warning( - "%d source tensors not found in target: %s...", - len(unmatched_source), unmatched_source[:3], - ) - - # For dtype-mismatched tensors, allocate temp buffers at source dtype - dtype_map = {"torch.bfloat16": torch.bfloat16, "torch.float16": torch.float16, - "torch.float32": torch.float32, "torch.uint8": torch.uint8, - "torch.float8_e4m3fn": torch.float8_e4m3fn} - cast_buffers = {} - for src_name, src_desc, dst_param, src_dtype_str in dtype_cast_needed: - src_torch_dtype = dtype_map.get(src_dtype_str, torch.bfloat16) - buf = torch.empty(dst_param.numel(), dtype=src_torch_dtype, device=f"cuda:{device_id}") - cast_buffers[src_name] = (buf, dst_param) - matched.append((src_name, src_desc, buf)) - - logger.info( - "Matched %d/%d params for direct RDMA transfer (%d need dtype cast)", - len(matched), len(source_descs), len(dtype_cast_needed), - ) - - # 5. Initialize NIXL and register TARGET param buffers - nixl_mgr = NixlTransferManager( - agent_name=f"trtllm-live-target-rank{mpi_rank}-{os.getpid()}", - device_id=device_id, - ) - nixl_mgr.initialize() - - # Register target params with NIXL (includes temp cast buffers) - dst_tensors = {name: param for name, _, param in matched} - nixl_mgr.register_tensors(dst_tensors) - - # 6. Build source descriptors for NIXL transfer - src_descs_for_transfer = [ - TensorDescriptor( - name=name, - addr=src_desc.addr, - size=src_desc.size, - device_id=src_desc.device_id, - dtype=src_desc.dtype, - ) - for name, src_desc, _ in matched - ] - - # 7. RDMA transfer: source params → target params - xfer_timeout = int(os.environ.get("MX_TRANSFER_TIMEOUT", "900")) - t0 = time.perf_counter() - bytes_transferred, n_tensors, _ = nixl_mgr.receive_from_source( - source_metadata=source_worker.nixl_metadata, - source_tensors=src_descs_for_transfer, - timeout_seconds=xfer_timeout, - ) - elapsed = time.perf_counter() - t0 - bw = (bytes_transferred * 8) / (elapsed * 1e9) if elapsed > 0 else 0 - - logger.info( - "Rank %d: transferred %d params (%.2f GB) in %.2fs (%.1f Gbps) — DIRECT into model params", - mpi_rank, n_tensors, bytes_transferred / 1e9, elapsed, bw, - ) - - # Diagnostic: checksum first few params to verify RDMA data - torch.cuda.synchronize(device_id) - for name, _, dst_param in matched[:5]: - val = dst_param.to(torch.float32) - cksum = val.sum().item() - nonzero = (dst_param != 0).sum().item() - logger.info( - "CHECKSUM rank %d: %s shape=%s dtype=%s sum=%.4f nonzero=%d/%d", - mpi_rank, name, list(dst_param.shape), dst_param.dtype, - cksum, nonzero, dst_param.numel(), - ) - - # 7.5. Apply dtype casts for mismatched tensors - for src_name, (buf, dst_param) in cast_buffers.items(): - dst_param.data.copy_(buf.to(dst_param.dtype)) - logger.info("Cast %s: %s → %s", src_name, buf.dtype, dst_param.dtype) - if cast_buffers: - logger.info("Applied %d dtype casts", len(cast_buffers)) - - nixl_mgr.shutdown() - - # 8. Load any size-mismatched tensors from PVC checkpoint as fallback - fallback_weights = {} - size_mismatched = { - src_name for src_name, src_desc in source_descs.items() - if src_name in target_params - and src_desc.size != target_params[src_name].numel() * target_params[src_name].element_size() - } - if size_mismatched: - logger.info( - "Loading %d size-mismatched tensors from PVC fallback: %s...", - len(size_mismatched), list(size_mismatched)[:3], - ) - try: - from safetensors import safe_open - import glob as _glob - safetensor_files = sorted(_glob.glob(os.path.join(checkpoint_dir, "*.safetensors"))) - for sf_path in safetensor_files: - with safe_open(sf_path, framework="pt", device=f"cuda:{device_id}") as f: - for key in f.keys(): - if key in size_mismatched: - fallback_weights[key] = f.get_tensor(key) - size_mismatched.discard(key) - if not size_mismatched: - break - if size_mismatched: - logger.warning("Still missing after PVC fallback: %s", size_mismatched) - except Exception as e: - logger.warning("PVC fallback failed: %s", e) - - # Return fallback weights for TRT-LLM to apply; P2P weights are already in model params - return fallback_weights - - def cleanup(self): - pass - - def _query_source(self, mx_server, model_name, timeout=600): - import grpc - - identity = _build_trtllm_identity(model_name=model_name) - mx_client = MxClient(server_url=mx_server) - - start = time.time() - while time.time() - start < timeout: - try: - list_resp = mx_client.list_sources( - identity=identity, - ) - if list_resp.instances: - workers = [] - for inst in list_resp.instances: - meta_resp = mx_client.get_metadata( - mx_source_id=inst.mx_source_id, - worker_id=inst.worker_id, - ) - if meta_resp.found and meta_resp.worker.tensors: - workers.append(meta_resp.worker) - - if workers: - logger.info("Found source: %d workers", len(workers)) - - class _SourceMeta: - pass - - result = _SourceMeta() - result.workers = workers - self._source_meta = result - mx_client.close() - return result - except grpc.RpcError as e: - logger.warning("Query failed: %s", e) - time.sleep(5) - - mx_client.close() - raise TimeoutError(f"Source for '{model_name}' not found after {timeout}s") - - -def _import_trtllm_for_config(): - from tensorrt_llm._torch.models.checkpoints.hf.config_loader import ( - HfConfigLoader, - ) - - return {"HfConfigLoader": HfConfigLoader} - - -class MxConfigLoader: - def load(self, checkpoint_dir: str, **kwargs): - trtllm = _import_trtllm_for_config() - HfConfigLoader = trtllm["HfConfigLoader"] - - logger.info("Loading config from local path: %s", checkpoint_dir) - return HfConfigLoader().load(checkpoint_dir, **kwargs) - - def cleanup(self): - pass - - -class MxLiveCheckpointLoader: - """ - Checkpoint loader that uses MxLiveWeightLoader for direct param-to-param transfer. - - Combines MxConfigLoader (config from MX server) with MxLiveWeightLoader - (direct RDMA into model params). - """ - - def __init__(self, mx_server: Optional[str] = None): - # Pass mx_server to weight loader so it's available even if env var isn't set - # when load_weights() is called in a different process context - self._weight_loader = MxLiveWeightLoader(mx_server=mx_server) - self._config_loader = None # Lazy init - self._weight_mapper = None - self._checkpoint_format = "mx-p2p" - - def get_default_weight_loader(self): - return MxLiveWeightLoader() - - def get_default_config_loader(self): - return MxConfigLoader() - - def cleanup(self): - if self._weight_mapper is not None and hasattr(self._weight_mapper, 'cleanup'): - self._weight_mapper.cleanup() - if self._weight_loader is not None: - self._weight_loader.cleanup() - - @property - def weight_loader(self): - return self._weight_loader - - @property - def weight_mapper(self): - return self._weight_mapper - - @weight_mapper.setter - def weight_mapper(self, value): - self._weight_mapper = value - - @property - def config_loader(self): - if self._config_loader is None: - self._config_loader = self.get_default_config_loader() - return self._config_loader - - @property - def checkpoint_format(self): - return self._checkpoint_format - - def load_config(self, checkpoint_dir: str, **kwargs): - logger.info("MxLiveCheckpointLoader.load_config(%s)", checkpoint_dir) - return self.config_loader.load(checkpoint_dir, **kwargs) - - def load_weights(self, checkpoint_dir: str, mapping=None, model=None, **kwargs): - logger.info("MxLiveCheckpointLoader.load_weights(model=%s)", model is not None) - return self._weight_loader.load_weights( - checkpoint_dir, mapping=mapping, model=model, **kwargs - ) - - def get_initialized_weight_mapper(self, model, config): - from tensorrt_llm._torch.models.checkpoints.auto_mapper import AutoCheckpointMapper - - if config.pretrained_config and config.pretrained_config.architectures: - model_arch = config.pretrained_config.architectures[0] - else: - raise ValueError("Cannot determine model architecture from config") - - weight_mapper = AutoCheckpointMapper.get("HF", model_arch) - weight_mapper.init_model_and_config(model, config) - self._weight_mapper = weight_mapper - return weight_mapper diff --git a/modelexpress_client/python/tests/test_k8s_service_client.py b/modelexpress_client/python/tests/test_k8s_service_client.py index 7b06f2434..f383780cc 100644 --- a/modelexpress_client/python/tests/test_k8s_service_client.py +++ b/modelexpress_client/python/tests/test_k8s_service_client.py @@ -65,6 +65,26 @@ def test_factory_unknown_backend_raises(monkeypatch): create_metadata_client() +def test_resolve_metadata_server_url_uses_explicit_before_env(monkeypatch): + from modelexpress.metadata.client_factory import resolve_metadata_server_url + + monkeypatch.setenv("MODEL_EXPRESS_URL", "mx-from-env:8001") + + assert resolve_metadata_server_url("https://mx-explicit:9000") == "mx-explicit:9000" + + +def test_resolve_metadata_server_url_uses_env_without_default(monkeypatch): + from modelexpress.metadata.client_factory import resolve_metadata_server_url + + monkeypatch.delenv("MODEL_EXPRESS_URL", raising=False) + monkeypatch.delenv("MX_SERVER_ADDRESS", raising=False) + + assert resolve_metadata_server_url() is None + + monkeypatch.setenv("MX_SERVER_ADDRESS", "http://mx-from-env:8001") + assert resolve_metadata_server_url() == "mx-from-env:8001" + + # --------------------------------------------------------------------------- # publish_metadata / list_sources / update_status behavior # --------------------------------------------------------------------------- diff --git a/modelexpress_client/python/tests/test_model_streamer_strategy.py b/modelexpress_client/python/tests/test_model_streamer_strategy.py index 2f3e90011..5317f8dd4 100644 --- a/modelexpress_client/python/tests/test_model_streamer_strategy.py +++ b/modelexpress_client/python/tests/test_model_streamer_strategy.py @@ -77,53 +77,58 @@ def _make_strategy(self): return ModelStreamerStrategy() def test_available_with_s3_uri(self): - ctx = _make_load_context() + ctx = _make_load_context(model_streamer_uri="s3://bucket/model") strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "s3://bucket/model"}): - with patch("importlib.util.find_spec", return_value=MagicMock()): - assert strategy.is_available(ctx) is True + with patch("importlib.util.find_spec", return_value=MagicMock()): + assert strategy.is_available(ctx) is True def test_available_with_local_path_env(self): - ctx = _make_load_context() + ctx = _make_load_context(model_streamer_uri="/models/deepseek") strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "/models/deepseek"}): - with patch("importlib.util.find_spec", return_value=MagicMock()): - assert strategy.is_available(ctx) is True + with patch("importlib.util.find_spec", return_value=MagicMock()): + assert strategy.is_available(ctx) is True def test_available_with_gcs_uri(self): - ctx = _make_load_context() + ctx = _make_load_context(model_streamer_uri="gs://bucket/model") strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "gs://bucket/model"}): - with patch("importlib.util.find_spec", return_value=MagicMock()): - assert strategy.is_available(ctx) is True + with patch("importlib.util.find_spec", return_value=MagicMock()): + assert strategy.is_available(ctx) is True def test_available_with_local_path(self): - ctx = _make_load_context() + ctx = _make_load_context(model_streamer_uri="/models/deepseek-ai/DeepSeek-V3") strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "/models/deepseek-ai/DeepSeek-V3"}): - with patch("importlib.util.find_spec", return_value=MagicMock()): - assert strategy.is_available(ctx) is True + with patch("importlib.util.find_spec", return_value=MagicMock()): + assert strategy.is_available(ctx) is True def test_available_with_hf_model_id(self): - ctx = _make_load_context() + ctx = _make_load_context(model_streamer_uri="deepseek-ai/DeepSeek-V3") strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "deepseek-ai/DeepSeek-V3"}): - with patch("importlib.util.find_spec", return_value=MagicMock()): - assert strategy.is_available(ctx) is True + with patch("importlib.util.find_spec", return_value=MagicMock()): + assert strategy.is_available(ctx) is True - def test_unavailable_no_env(self): + def test_unavailable_no_context_uri(self): ctx = _make_load_context() strategy = self._make_strategy() - with patch.dict("os.environ", {}, clear=True): - with patch("importlib.util.find_spec", return_value=MagicMock()): - assert strategy.is_available(ctx) is False + with patch("importlib.util.find_spec", return_value=MagicMock()): + assert strategy.is_available(ctx) is False def test_unavailable_no_package(self): - ctx = _make_load_context() + ctx = _make_load_context(model_streamer_uri="s3://bucket/model") strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "s3://bucket/model"}): - with patch("importlib.util.find_spec", return_value=None): - assert strategy.is_available(ctx) is False + with patch("importlib.util.find_spec", return_value=None): + assert strategy.is_available(ctx) is False + + +def test_resolve_model_streamer_uri_keeps_mx_model_uri_as_gate(monkeypatch): + from modelexpress.load_strategy.context import resolve_model_streamer_uri + + model_config = SimpleNamespace( + model="deepseek-ai/DeepSeek-V3", + model_weights="s3://bucket/deepseek-ai/DeepSeek-V3", + ) + monkeypatch.setenv("MX_MODEL_URI", "1") + + assert resolve_model_streamer_uri(model_config) == model_config.model_weights # --------------------------------------------------------------------------- @@ -136,16 +141,25 @@ def _make_strategy(self): from modelexpress.load_strategy.model_streamer_strategy import ModelStreamerStrategy return ModelStreamerStrategy() - def _make_ctx_with_uri(self, *, model_weights=None, model="ignored"): + def _make_ctx_with_uri( + self, + *, + model_streamer_uri="s3://bucket/model", + model_weights=None, + model="ignored", + ): model_config = MagicMock() model_config.model_weights = model_weights model_config.model = model - return _make_load_context(model_config=model_config) + return _make_load_context( + model_config=model_config, + model_streamer_uri=model_streamer_uri, + ) @patch("modelexpress.load_strategy.model_streamer_strategy.register_tensors") def test_success_path_s3(self, mock_register): model = MagicMock() - ctx = self._make_ctx_with_uri(model_weights="s3://bucket/model") + ctx = self._make_ctx_with_uri(model_streamer_uri="s3://bucket/model") strategy = self._make_strategy() with patch( @@ -165,10 +179,10 @@ def test_success_path_s3(self, mock_register): mock_register.assert_called_once_with(result, ctx) @patch("modelexpress.load_strategy.model_streamer_strategy.register_tensors") - def test_success_path_local_falls_back_to_model(self, mock_register): - """When model_weights is unset, the URI comes from model_config.model.""" + def test_success_path_local_uses_context_uri(self, mock_register): + """The streaming URI comes from LoadContext.""" model = MagicMock() - ctx = self._make_ctx_with_uri(model_weights=None, model="/models/llama") + ctx = self._make_ctx_with_uri(model_streamer_uri="/models/llama") strategy = self._make_strategy() with patch( @@ -184,29 +198,30 @@ def test_success_path_local_falls_back_to_model(self, mock_register): mock_stream.assert_called_once_with("/models/llama", ctx, model) @patch("modelexpress.load_strategy.model_streamer_strategy.register_tensors") - def test_uri_from_model_weights_not_from_env(self, mock_register): - """The streaming URI comes from model_config, not from MX_MODEL_URI.""" + def test_context_uri_wins_over_model_config(self, mock_register): + """LoadContext keeps loader config immutable after context creation.""" model = MagicMock() ctx = self._make_ctx_with_uri( - model_weights="s3://bucket/from-config", model="/ignored" + model_streamer_uri="s3://bucket/from-context", + model_weights="s3://bucket/from-config", + model="/ignored", ) strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "s3://other/path-in-env"}): - with patch( - "modelexpress.load_strategy.model_streamer_strategy." - "ModelStreamerStrategy._stream_weights" - ) as mock_stream: - mock_stream.side_effect = RuntimeError("expected") - with pytest.raises(StrategyFailed, match="expected"): - strategy.load(model, ctx) + with patch( + "modelexpress.load_strategy.model_streamer_strategy." + "ModelStreamerStrategy._stream_weights" + ) as mock_stream: + mock_stream.side_effect = RuntimeError("expected") + with pytest.raises(StrategyFailed, match="expected"): + strategy.load(model, ctx) - mock_stream.assert_called_once_with("s3://bucket/from-config", ctx, model) + mock_stream.assert_called_once_with("s3://bucket/from-context", ctx, model) @patch("modelexpress.load_strategy.model_streamer_strategy.register_tensors") def test_raises_strategy_failed_on_error(self, mock_register): model = MagicMock() - ctx = self._make_ctx_with_uri(model_weights="s3://bucket/model") + ctx = self._make_ctx_with_uri(model_streamer_uri="s3://bucket/model") strategy = self._make_strategy() with patch( @@ -225,17 +240,19 @@ def test_apply_weight_iter_failure_is_mutated(self, mock_register): model = MagicMock() adapter = _FakeAdapter() adapter.apply_weight_iter = MagicMock(side_effect=RuntimeError("partial load")) - ctx = _make_load_context(adapter=adapter) + ctx = _make_load_context( + adapter=adapter, + model_streamer_uri="s3://bucket/model", + ) strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "s3://bucket/model"}): - with patch( - "modelexpress.load_strategy.model_streamer_strategy." - "ModelStreamerStrategy._stream_weights", - return_value=iter([("layer.0.weight", torch.randn(4, 4))]), - ): - with pytest.raises(StrategyFailed, match="partial load") as exc: - strategy.load(model, ctx) + with patch( + "modelexpress.load_strategy.model_streamer_strategy." + "ModelStreamerStrategy._stream_weights", + return_value=iter([("layer.0.weight", torch.randn(4, 4))]), + ): + with pytest.raises(StrategyFailed, match="partial load") as exc: + strategy.load(model, ctx) assert exc.value.mutated is True mock_register.assert_not_called() @@ -245,17 +262,19 @@ def test_after_weight_iter_failure_is_mutated(self, mock_register): model = MagicMock() adapter = _FakeAdapter() adapter.after_weight_iter_load = MagicMock(side_effect=RuntimeError("post load")) - ctx = _make_load_context(adapter=adapter) + ctx = _make_load_context( + adapter=adapter, + model_streamer_uri="s3://bucket/model", + ) strategy = self._make_strategy() - with patch.dict("os.environ", {"MX_MODEL_URI": "s3://bucket/model"}): - with patch( - "modelexpress.load_strategy.model_streamer_strategy." - "ModelStreamerStrategy._stream_weights", - return_value=iter([("layer.0.weight", torch.randn(4, 4))]), - ): - with pytest.raises(StrategyFailed, match="post load") as exc: - strategy.load(model, ctx) + with patch( + "modelexpress.load_strategy.model_streamer_strategy." + "ModelStreamerStrategy._stream_weights", + return_value=iter([("layer.0.weight", torch.randn(4, 4))]), + ): + with pytest.raises(StrategyFailed, match="post load") as exc: + strategy.load(model, ctx) assert exc.value.mutated is True mock_register.assert_not_called() @@ -318,10 +337,11 @@ def _patch_runai_loader(self, tensors): def test_uses_vllm_native_iterator_and_enables_distributed(self): tensor = torch.randn(2, 2) - adapter, _load_config = self._make_adapter(tp_size=8) + with patch.dict("os.environ", {"MX_MS_DISTRIBUTED": "1"}): + adapter, _load_config = self._make_adapter(tp_size=8) patcher, loader_cls, loader_instance = self._patch_runai_loader([("w", tensor)]) - with patcher, patch.dict("os.environ", {"MX_MS_DISTRIBUTED": "1"}): + with patcher: weights = list(adapter.build_model_streamer_weight_iter("az://models/model")) native_load_config = loader_cls.call_args.args[0] @@ -333,26 +353,28 @@ def test_uses_vllm_native_iterator_and_enables_distributed(self): assert weights[0][1] is tensor def test_preserves_existing_extra_config_when_distributed_disabled(self): - adapter, _load_config = self._make_adapter( - tp_size=1, - extra_config={"concurrency": 4}, - ) + with patch.dict("os.environ", {"MX_MS_DISTRIBUTED": "1"}): + adapter, _load_config = self._make_adapter( + tp_size=1, + extra_config={"concurrency": 4}, + ) patcher, loader_cls, _loader_instance = self._patch_runai_loader([]) - with patcher, patch.dict("os.environ", {"MX_MS_DISTRIBUTED": "1"}): + with patcher: list(adapter.build_model_streamer_weight_iter("az://models/model")) native_load_config = loader_cls.call_args.args[0] assert native_load_config.model_loader_extra_config == {"concurrency": 4} def test_distributed_disabled_by_default_even_with_tp_gt_one(self): - adapter, _load_config = self._make_adapter( - tp_size=8, - extra_config={"concurrency": 4}, - ) + with patch.dict("os.environ", {}, clear=True): + adapter, _load_config = self._make_adapter( + tp_size=8, + extra_config={"concurrency": 4}, + ) patcher, loader_cls, _loader_instance = self._patch_runai_loader([]) - with patcher, patch.dict("os.environ", {}, clear=True): + with patcher: list(adapter.build_model_streamer_weight_iter("az://models/model")) native_load_config = loader_cls.call_args.args[0] diff --git a/modelexpress_client/python/tests/test_sglang_loader.py b/modelexpress_client/python/tests/test_sglang_loader.py index 700573c67..4b947eff7 100644 --- a/modelexpress_client/python/tests/test_sglang_loader.py +++ b/modelexpress_client/python/tests/test_sglang_loader.py @@ -253,18 +253,17 @@ def test_sglang_adapter_enables_distributed_model_streamer(monkeypatch): load_format = SimpleNamespace(RUNAI_STREAMER="runai_streamer") _install_sglang_runai_loader_modules(monkeypatch, loader_cls, load_format) - adapter = SglangAdapter( - _load_config(model_loader_extra_config={"concurrency": 4}), - _model_config(), - _device_config(device="cuda:0", gpu_id=0), - ) + with patch.dict("os.environ", {"MX_MS_DISTRIBUTED": "1"}): + adapter = SglangAdapter( + _load_config(model_loader_extra_config={"concurrency": 4}), + _model_config(), + _device_config(device="cuda:0", gpu_id=0), + ) with patch( "modelexpress.engines.sglang.adapter._get_parallel_size", return_value=8, - ), patch.object(adapter, "is_cuda_alike", return_value=True), patch.dict( - "os.environ", {"MX_MS_DISTRIBUTED": "1"} - ): + ), patch.object(adapter, "is_cuda_alike", return_value=True): list( adapter.build_model_streamer_weight_iter( "s3://bucket/deepseek-ai/DeepSeek-V3", diff --git a/modelexpress_client/python/tests/test_trtllm_loader.py b/modelexpress_client/python/tests/test_trtllm_loader.py new file mode 100644 index 000000000..bbcdf441c --- /dev/null +++ b/modelexpress_client/python/tests/test_trtllm_loader.py @@ -0,0 +1,882 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for the TensorRT-LLM ModelExpress loader integration.""" + +from __future__ import annotations + +import importlib +import sys +import types +from contextlib import contextmanager +from types import SimpleNamespace + +import pytest + +from modelexpress import p2p_pb2 +from modelexpress.adapter import EngineAdapter +from modelexpress.engines.trtllm import TrtllmAdapter +from modelexpress.engines.trtllm import adapter as trtllm_adapter +from modelexpress.engines.trtllm import loader as trtllm +from modelexpress.load_strategy.context import LoadResult + + +class _FakeTensor: + def numel(self): + return 2 + + def element_size(self): + return 4 + + +def _fake_trtllm_mapping(rank=0, tp_size=1, pp_size=1, cp_size=1): + return SimpleNamespace( + rank=rank, + tp_size=tp_size, + pp_size=pp_size, + cp_size=cp_size, + ) + + +@contextmanager +def _noop_rank_log_scope(global_rank): + del global_rank + yield + + +@pytest.fixture +def trtllm_loader_with_fake_trt(monkeypatch): + """Reload the TRT-LLM loader with a minimal fake TensorRT-LLM surface.""" + + def install_package(name): + module = types.ModuleType(name) + module.__path__ = [] + monkeypatch.setitem(sys.modules, name, module) + parent, _, child = name.rpartition(".") + if parent: + setattr(sys.modules[parent], child, module) + return module + + def install_module(name): + module = types.ModuleType(name) + monkeypatch.setitem(sys.modules, name, module) + parent, _, child = name.rpartition(".") + setattr(sys.modules[parent], child, module) + return module + + for package in [ + "tensorrt_llm", + "tensorrt_llm._torch", + "tensorrt_llm._torch.models", + "tensorrt_llm._torch.models.checkpoints", + "tensorrt_llm._torch.models.checkpoints.hf", + ]: + install_package(package) + + base_config = install_module( + "tensorrt_llm._torch.models.checkpoints.base_config_loader" + ) + base_weight = install_module( + "tensorrt_llm._torch.models.checkpoints.base_weight_loader" + ) + base_mapper = install_module( + "tensorrt_llm._torch.models.checkpoints.base_weight_mapper" + ) + auto_mapper = install_module( + "tensorrt_llm._torch.models.checkpoints.auto_mapper" + ) + hf_loader = install_module( + "tensorrt_llm._torch.models.checkpoints.hf.checkpoint_loader" + ) + modeling_utils = install_module("tensorrt_llm._torch.models.modeling_utils") + mapping_module = install_module("tensorrt_llm.mapping") + + class BaseConfigLoader: + pass + + class BaseWeightLoader: + pass + + class ConsumableWeightsDict: + def __init__(self, weights): + self._weights = weights + + def __len__(self): + return len(self._weights) + + def values(self): + return list(self._weights.values()) + + class BaseWeightMapper: + pass + + class FakeWeightMapper: + def __init__(self): + self.init_calls = [] + + def init_model_and_config(self, model, config): + self.init_calls.append((model, config)) + + class AutoCheckpointMapper: + calls = [] + + @staticmethod + def get(format, name=None): + AutoCheckpointMapper.calls.append((format, name)) + return FakeWeightMapper() + + class HfCheckpointLoader: + def __init__(self, *, weight_loader=None, weight_mapper=None, config_loader=None): + self.weight_loader = weight_loader + self.weight_mapper = weight_mapper + self.config_loader = config_loader + self.disk_loads = [] + + def load_weights(self, checkpoint_dir, *, mapping, **kwargs): + self.disk_loads.append((checkpoint_dir, mapping, kwargs)) + return {"disk.weight": _FakeTensor()} + + class Mapping: + pass + + def register_checkpoint_loader(name): + def decorator(cls): + cls._registered_checkpoint_loader_name = name + return cls + + return decorator + + base_config.BaseConfigLoader = BaseConfigLoader + base_weight.BaseWeightLoader = BaseWeightLoader + base_weight.ConsumableWeightsDict = ConsumableWeightsDict + base_mapper.BaseWeightMapper = BaseWeightMapper + auto_mapper.AutoCheckpointMapper = AutoCheckpointMapper + hf_loader.HfCheckpointLoader = HfCheckpointLoader + modeling_utils.register_checkpoint_loader = register_checkpoint_loader + mapping_module.Mapping = Mapping + + reloaded = importlib.reload(trtllm) + monkeypatch.setattr(reloaded, "_rank_log_scope", _noop_rank_log_scope) + yield reloaded + importlib.reload(trtllm) + + +def test_trtllm_adapter_inherits_engine_adapter(): + assert issubclass(TrtllmAdapter, EngineAdapter) + + +def test_trtllm_adapter_native_loader_feeds_default_strategy(): + fallback_weights = {"disk.weight": _FakeTensor()} + model = object() + adapter = TrtllmAdapter( + model_name="Qwen/Qwen2.5-7B", + model=model, + native_loader=lambda: fallback_weights, + ) + + result = adapter.load_via_native(LoadResult(value=model, model=model)) + + assert result.value is fallback_weights + assert result.model is None + assert result.publishable is False + + +def test_trtllm_rdma_receiver_is_not_republished_by_chain(): + model = object() + adapter = TrtllmAdapter(model_name="Qwen/Qwen2.5-7B", model=model) + result = LoadResult(value=model, model=model) + + result = adapter.after_rdma_receive(result) + + assert result.publishable is False + + +def test_trtllm_identity_is_internal_to_adapter(): + adapter = TrtllmAdapter(model_name="Qwen/Qwen2.5-7B") + identity = adapter.build_identity() + + assert identity.model_name == "Qwen/Qwen2.5-7B" + assert identity.backend_framework == p2p_pb2.BACKEND_FRAMEWORK_TRT_LLM + assert identity.dtype == "unknown" + + +def test_trtllm_identity_uses_live_parameter_dtype(): + model = SimpleNamespace( + parameters=lambda: iter([SimpleNamespace(dtype=trtllm_adapter.torch.float16)]) + ) + adapter = TrtllmAdapter(model_name="Qwen/Qwen2.5-7B", model=model) + identity = adapter.build_identity() + + assert identity.dtype == "float16" + + +def test_trtllm_load_context_uses_model_config_dtype_and_quantization(monkeypatch): + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + model = SimpleNamespace( + model_config=SimpleNamespace( + pretrained_config=SimpleNamespace( + torch_dtype=trtllm_adapter.torch.float16, + model_type="qwen2", + ), + quant_config=SimpleNamespace(quant_algo="nvfp4"), + ), + parameters=lambda: iter(()), + ) + + ctx = trtllm_adapter.build_trtllm_load_context( + model_name="Qwen/Qwen2.5-7B", + checkpoint_dir="/models/qwen", + model=model, + ) + + assert ctx.model_config.dtype is trtllm_adapter.torch.float16 + assert ctx.model_config.quantization == "nvfp4" + assert ctx.model_config.hf_text_config.model_type == "qwen2" + assert ctx.identity.dtype == "float16" + assert ctx.identity.quantization == "nvfp4" + + +def test_trtllm_worker_rank_is_shard_key_not_global_rank(monkeypatch): + del monkeypatch + mapping = SimpleNamespace( + rank=5, + tp_size=2, + pp_size=2, + cp_size=2, + tp_rank=0, + pp_rank=1, + moe_ep_size=1, + ) + + adapter = TrtllmAdapter(model_name="Qwen/Qwen2.5-7B", mapping=mapping) + identity = adapter.build_identity() + + assert adapter.get_global_rank() == 5 + assert adapter.get_worker_rank() == 2 + assert identity.tensor_parallel_size == 2 + assert identity.pipeline_parallel_size == 2 + + +def test_trtllm_worker_rank_collapses_context_parallel_replicas(monkeypatch): + # TRT-LLM CP ranks share the same model weights for a PP/TP shard, so they + # must match to the same source worker. + del monkeypatch + cp_rank_0 = SimpleNamespace( + rank=2, + tp_size=2, + pp_size=2, + cp_size=2, + tp_rank=1, + pp_rank=0, + cp_rank=0, + ) + cp_rank_1 = SimpleNamespace( + rank=3, + tp_size=2, + pp_size=2, + cp_size=2, + tp_rank=1, + pp_rank=0, + cp_rank=1, + ) + + assert ( + TrtllmAdapter( + model_name="Qwen/Qwen2.5-7B", + mapping=cp_rank_0, + ).get_worker_rank() + == 1 + ) + assert ( + TrtllmAdapter( + model_name="Qwen/Qwen2.5-7B", + mapping=cp_rank_1, + ).get_worker_rank() + == 1 + ) + + +def test_checkpoint_loader_runs_shared_load_strategy_chain( + monkeypatch, + trtllm_loader_with_fake_trt, +): + calls = {} + + class FakeClient: + pass + + def fake_create_metadata_client(**kwargs): + calls["client_kwargs"] = kwargs + return FakeClient() + + def fake_chain_run(model_arg, ctx): + calls["chain_model"] = model_arg + calls["ctx"] = ctx + ctx.selected_strategy = "rdma" + return model_arg + + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + fake_create_metadata_client, + ) + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + mapping = _fake_trtllm_mapping() + model = object() + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + model_name="Qwen/Qwen2.5-7B", + ) + + result = loader.load_weights("/models/qwen", mapping=mapping, model=model) + + assert result == {} + assert loader.p2p_succeeded is True + assert calls["chain_model"] is model + assert calls["ctx"].adapter.mapping is mapping + assert calls["ctx"].adapter.model_name == "Qwen/Qwen2.5-7B" + assert calls["ctx"].identity.model_name == "Qwen/Qwen2.5-7B" + assert calls["client_kwargs"] == { + "worker_rank": 0, + "server_url": "mx.example:8001", + } + assert calls["ctx"].metadata_server_url == "mx.example:8001" + + +def test_checkpoint_loader_uses_modelexpress_load_format(trtllm_loader_with_fake_trt): + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader() + + assert loader.checkpoint_format == "modelexpress" + + +def test_checkpoint_loader_returns_default_strategy_weights( + monkeypatch, + trtllm_loader_with_fake_trt, +): + fallback_weights = {"disk.weight": _FakeTensor()} + + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + ctx.selected_strategy = "default" + return fallback_weights + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + ) + + result = loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=object(), + ) + + assert result is fallback_weights + assert loader.p2p_succeeded is False + + +def test_checkpoint_loader_default_strategy_none_is_not_p2p( + monkeypatch, + trtllm_loader_with_fake_trt, +): + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + ctx.selected_strategy = "default" + return None + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + ) + + result = loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=object(), + ) + + assert result == {} + assert loader.p2p_succeeded is False + + +def test_checkpoint_loader_preserves_native_consumable_weights( + monkeypatch, + trtllm_loader_with_fake_trt, +): + fallback_weights = trtllm_loader_with_fake_trt.ConsumableWeightsDict( + {"disk.weight": _FakeTensor()} + ) + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + ctx.selected_strategy = "default" + return fallback_weights + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + ) + + result = loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=object(), + ) + + assert result is fallback_weights + assert loader.p2p_succeeded is False + + +def test_checkpoint_loader_uses_hf_mapper_for_mx_fallback( + trtllm_loader_with_fake_trt, +): + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader() + model = object() + config = SimpleNamespace( + pretrained_config=SimpleNamespace(architectures=["Qwen3ForCausalLM"]) + ) + + mapper = loader.get_initialized_weight_mapper(model, config) + + assert trtllm_loader_with_fake_trt.AutoCheckpointMapper.calls == [ + ("HF", "Qwen3ForCausalLM") + ] + assert mapper.init_calls == [(model, config)] + + +def test_checkpoint_loader_model_kwarg_is_not_forwarded_to_disk_fallback( + trtllm_loader_with_fake_trt, +): + mapping = object() + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader() + + result = loader.load_weights( + "/models/qwen", + mapping=mapping, + model=None, + extra_option="kept", + ) + + assert list(result) == ["disk.weight"] + assert loader.disk_loads == [ + ("/models/qwen", mapping, {"extra_option": "kept"}), + ] + + +def test_model_less_load_does_not_clear_publish_context( + monkeypatch, + trtllm_loader_with_fake_trt, +): + captured = {} + + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + captured["load_ctx"] = ctx + ctx.selected_strategy = "rdma" + return model + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + def fake_publish(result, ctx): + captured["publish_result"] = result + captured["publish_ctx"] = ctx + + import modelexpress.load_strategy as load_strategy + + monkeypatch.setattr(load_strategy, "publish_loaded_model", fake_publish) + + model = object() + mapping = _fake_trtllm_mapping() + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + model_name="Qwen/Qwen2.5-7B", + ) + + loader.load_weights("/models/qwen", mapping=mapping, model=model) + loader.load_weights("/models/draft", mapping=mapping) + loader.post_load_publish(model, checkpoint_dir="/models/qwen") + + assert captured["publish_result"].model is model + assert captured["publish_ctx"] is captured["load_ctx"] + assert loader.disk_loads == [("/models/draft", mapping, {})] + + +def test_post_load_publish_reuses_load_context_worker_id( + monkeypatch, + trtllm_loader_with_fake_trt, +): + captured = {} + fallback_weights = {"disk.weight": _FakeTensor()} + + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + captured["load_ctx"] = ctx + ctx.selected_strategy = "default" + return fallback_weights + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + def fake_publish(result, ctx): + captured["publish_result"] = result + captured["publish_ctx"] = ctx + + import modelexpress.load_strategy as load_strategy + + monkeypatch.setattr(load_strategy, "publish_loaded_model", fake_publish) + + model = object() + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + model_name="Qwen/Qwen2.5-7B", + ) + + result = loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=model, + ) + loader.post_load_publish(model, checkpoint_dir="/models/qwen") + + assert result is fallback_weights + assert captured["publish_result"].model is model + assert captured["publish_ctx"] is captured["load_ctx"] + assert captured["publish_ctx"].worker_id == captured["load_ctx"].worker_id + + +def test_post_load_publish_publishes_rdma_receiver( + monkeypatch, + trtllm_loader_with_fake_trt, +): + captured = {} + + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + captured["load_ctx"] = ctx + ctx.selected_strategy = "rdma" + return model + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + def fake_publish(result, ctx): + captured["publish_result"] = result + captured["publish_ctx"] = ctx + + import modelexpress.load_strategy as load_strategy + + monkeypatch.setattr(load_strategy, "publish_loaded_model", fake_publish) + + model = object() + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + model_name="Qwen/Qwen2.5-7B", + ) + + result = loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=model, + ) + loader.post_load_publish( + model, + checkpoint_dir="/models/qwen", + weights_preloaded=True, + ) + + assert result == {} + assert loader.p2p_succeeded is True + assert captured["publish_result"].model is model + assert captured["publish_ctx"] is captured["load_ctx"] + assert captured["publish_ctx"].worker_id == captured["load_ctx"].worker_id + + +def test_publish_as_source_reuses_matching_load_context( + monkeypatch, + trtllm_loader_with_fake_trt, +): + captured = {} + fallback_weights = {"disk.weight": _FakeTensor()} + + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + captured["load_ctx"] = ctx + ctx.selected_strategy = "default" + return fallback_weights + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + def fake_publish(result, ctx): + captured["publish_result"] = result + captured["publish_ctx"] = ctx + + import modelexpress.load_strategy as load_strategy + + monkeypatch.setattr(load_strategy, "publish_loaded_model", fake_publish) + + model = object() + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + model_name="Qwen/Qwen2.5-7B", + ) + build_calls = [] + original_build_context = loader._build_load_context + + def counting_build_context(**kwargs): + build_calls.append(kwargs) + return original_build_context(**kwargs) + + loader._build_load_context = counting_build_context + + loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=model, + ) + loader.publish_as_source(model, checkpoint_dir="/models/qwen") + + assert captured["publish_result"].model is model + assert captured["publish_ctx"] is captured["load_ctx"] + assert len(build_calls) == 1 + + +def test_checkpoint_loader_timeout_waits_then_uses_chain_default( + monkeypatch, + trtllm_loader_with_fake_trt, +): + calls = {} + + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + + def fake_chain_run(model, ctx): + calls["chain_ctx"] = ctx + ctx.selected_strategy = "default" + return ctx.adapter.load_via_native(LoadResult(value=model, model=model)).value + + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + fake_chain_run, + ) + + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + query_timeout_s=0, + ) + + result = loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=object(), + ) + + assert list(result) == ["disk.weight"] + assert calls["chain_ctx"].worker_rank == 0 + assert calls["chain_ctx"].load_config.source_query_timeout_s == 0 + assert loader.p2p_succeeded is False + + +def test_checkpoint_loader_falls_back_when_chain_fails( + monkeypatch, + trtllm_loader_with_fake_trt, +): + monkeypatch.setattr( + trtllm_adapter, + "create_metadata_client", + lambda **kwargs: SimpleNamespace(), + ) + monkeypatch.setattr( + trtllm_loader_with_fake_trt.LoadStrategyChain, + "run", + lambda model, ctx: (_ for _ in ()).throw(RuntimeError("transfer failed")), + ) + + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + ) + + result = loader.load_weights( + "/models/qwen", + mapping=_fake_trtllm_mapping(), + model=object(), + ) + + assert list(result) == ["disk.weight"] + assert loader.p2p_succeeded is False + + +def test_checkpoint_loader_publish_without_load_context_skips( + caplog, + monkeypatch, + trtllm_loader_with_fake_trt, +): + published = [] + + def fake_publish(model, ctx): + published.append((model, ctx)) + + monkeypatch.setattr("modelexpress.load_strategy.publish_loaded_model", fake_publish) + + model = object() + loader = trtllm_loader_with_fake_trt.MXCheckpointLoader( + mx_server_url="mx.example:8001", + model_name="Qwen/Qwen2.5-7B", + ) + + with caplog.at_level(trtllm_loader_with_fake_trt.logging.WARNING): + loader._publish_as_source(model, checkpoint_dir="/unused/path", mapping=object()) + + assert published == [] + assert "matching load context" in caplog.text + + +def test_publish_loaded_model_uses_shared_strategy_helpers(monkeypatch): + captured = {} + + nixl_manager = object() + ctx = SimpleNamespace(nixl_manager=None) + model = SimpleNamespace() + + def fake_register_tensors(result, ctx_arg): + captured["register_result"] = result + captured["register_ctx"] = ctx_arg + ctx_arg.nixl_manager = nixl_manager + + def fake_publish_metadata(ctx_arg): + captured["publish_ctx"] = ctx_arg + + import modelexpress.load_strategy.base as load_strategy_base + from modelexpress.load_strategy import publish_loaded_model + + monkeypatch.setattr(load_strategy_base, "register_tensors", fake_register_tensors) + monkeypatch.setattr(load_strategy_base, "publish_metadata", fake_publish_metadata) + + publish_loaded_model(LoadResult(value=model, model=model), ctx) + + assert captured["register_result"].model is model + assert captured["register_ctx"] is ctx + assert captured["publish_ctx"] is ctx + assert ctx.nixl_manager is nixl_manager + assert not hasattr(model, "_mx_nixl_managers") + assert model._mx_load_context is ctx + + +def test_publish_loaded_model_skips_non_publishable_result(monkeypatch): + captured = {} + ctx = SimpleNamespace() + model = SimpleNamespace() + result = LoadResult(value=model, model=model, publishable=False) + + def fake_register_tensors(result, ctx_arg): + captured["register"] = (result, ctx_arg) + + def fake_publish_metadata(ctx_arg): + captured["publish"] = ctx_arg + + import modelexpress.load_strategy.base as load_strategy_base + from modelexpress.load_strategy import publish_loaded_model + + monkeypatch.setattr(load_strategy_base, "register_tensors", fake_register_tensors) + monkeypatch.setattr(load_strategy_base, "publish_metadata", fake_publish_metadata) + + publish_loaded_model(result, ctx) + + assert captured == {} + assert not hasattr(model, "_mx_load_context") + + +def test_rank_log_scope_flushes_fsyncs_and_mirrors_to_stderr( + tmp_path, monkeypatch, capsys +): + fsync_calls = [] + monkeypatch.setenv("MX_TRANSFER_LOG_DIR", str(tmp_path)) + monkeypatch.setattr(trtllm.os, "fsync", lambda fd: fsync_calls.append(fd)) + + with trtllm._rank_log_scope(3): + trtllm.logging.getLogger("modelexpress").info("transfer metric") + with trtllm._rank_log_scope(3): + trtllm.logging.getLogger("modelexpress").info("publish metric") + + assert fsync_calls + rank_log_text = (tmp_path / "rank3.log").read_text() + stderr_text = capsys.readouterr().err + assert "[ModelExpress] [I] transfer metric" in rank_log_text + assert "[ModelExpress] [I] publish metric" in rank_log_text + assert "[ModelExpress] [I] transfer metric" in stderr_text + assert "[ModelExpress] [I] publish metric" in stderr_text diff --git a/modelexpress_client/python/tests/test_vllm_adapter.py b/modelexpress_client/python/tests/test_vllm_adapter.py index 2094cd292..db8669e4b 100644 --- a/modelexpress_client/python/tests/test_vllm_adapter.py +++ b/modelexpress_client/python/tests/test_vllm_adapter.py @@ -89,6 +89,40 @@ def test_build_vllm_load_context_keeps_explicit_cuda_index(monkeypatch): assert ctx.device_id == ctx.target_device.index +def test_build_vllm_load_context_populates_metadata_url_from_env(monkeypatch): + _stub_vllm_current_device(monkeypatch, current_device=0) + captured = {} + + def fake_create_metadata_client(worker_rank, server_url=None): + captured["worker_rank"] = worker_rank + captured["server_url"] = server_url + return object() + + monkeypatch.setattr( + "modelexpress.engines.vllm.adapter.create_metadata_client", + fake_create_metadata_client, + ) + monkeypatch.setenv("MODEL_EXPRESS_URL", "http://mx.example:8001") + vllm_config = _context_config(load_device=None) + + ctx = build_vllm_load_context(vllm_config, _model_config()) + + assert ctx.metadata_server_url == "mx.example:8001" + assert captured["server_url"] == "mx.example:8001" + + +def test_build_vllm_load_context_preserves_mx_model_uri_gate(monkeypatch): + _stub_vllm_current_device(monkeypatch, current_device=0) + _stub_metadata_client(monkeypatch) + monkeypatch.setenv("MX_MODEL_URI", "1") + model_config = _model_config() + model_config.model_weights = "s3://bucket/from-config" + + ctx = build_vllm_load_context(_context_config(load_device=None), model_config) + + assert ctx.model_streamer_uri == "s3://bucket/from-config" + + def _stub_vllm_current_device(monkeypatch, *, current_device: int) -> None: fake_platforms = SimpleNamespace( current_platform=SimpleNamespace( @@ -101,7 +135,7 @@ def _stub_vllm_current_device(monkeypatch, *, current_device: int) -> None: def _stub_metadata_client(monkeypatch) -> None: monkeypatch.setattr( "modelexpress.engines.vllm.adapter.create_metadata_client", - lambda worker_rank: object(), + lambda worker_rank, server_url=None: object(), ) diff --git a/modelexpress_client/python/tests/test_vllm_loader.py b/modelexpress_client/python/tests/test_vllm_loader.py index 16acde0ef..f4a311799 100644 --- a/modelexpress_client/python/tests/test_vllm_loader.py +++ b/modelexpress_client/python/tests/test_vllm_loader.py @@ -410,19 +410,17 @@ def test_skips_when_no_metadata_path_configured(self, mock_available): @patch("modelexpress.load_strategy.base.is_nixl_available", return_value=False) def test_skips_when_nixl_unavailable(self, _mock): from modelexpress.load_strategy.base import register_tensors - ctx = _make_load_context() + ctx = _make_load_context(metadata_server_url="localhost:8001") model = MagicMock() - with patch.dict(os.environ, {"MX_SERVER_ADDRESS": "localhost:8001"}): - register_tensors(model, ctx) + register_tensors(model, ctx) assert ctx.nixl_manager is None @patch("modelexpress.load_strategy.base.is_nixl_available", return_value=True) def test_requires_adapter_when_nixl_available(self, _mock): from modelexpress.load_strategy.base import register_tensors - ctx = _make_load_context(adapter=None) - with patch.dict(os.environ, {"MX_SERVER_ADDRESS": "localhost:8001"}): - with pytest.raises(RuntimeError, match="engine adapter"): - register_tensors(MagicMock(), ctx) + ctx = _make_load_context(adapter=None, metadata_server_url="localhost:8001") + with pytest.raises(RuntimeError, match="engine adapter"): + register_tensors(MagicMock(), ctx) @patch("modelexpress.load_strategy.base.is_nixl_available", return_value=True) @patch( @@ -431,10 +429,9 @@ def test_requires_adapter_when_nixl_available(self, _mock): ) def test_nixl_init_failure_does_not_raise(self, _init, _avail): from modelexpress.load_strategy.base import register_tensors - ctx = _make_load_context() + ctx = _make_load_context(metadata_server_url="localhost:8001") model = MagicMock() - with patch.dict(os.environ, {"MX_SERVER_ADDRESS": "localhost:8001"}): - register_tensors(model, ctx) + register_tensors(model, ctx) assert ctx.nixl_manager is None @patch("modelexpress.load_strategy.base.is_nixl_available", return_value=True) @@ -445,11 +442,10 @@ def test_tensor_registration_failure_does_not_raise(self, mock_init, _avail): mock_mgr.tensor_descriptors = [] mock_mgr.register_tensors.side_effect = RuntimeError("memory registration failed") mock_init.return_value = mock_mgr - ctx = _make_load_context() + ctx = _make_load_context(metadata_server_url="localhost:8001") ctx.adapter.discover_tensors = MagicMock(return_value={"t": MagicMock()}) model = MagicMock() - with patch.dict(os.environ, {"MX_SERVER_ADDRESS": "localhost:8001"}): - register_tensors(model, ctx) + register_tensors(model, ctx) class TestPublishMetadataErrorHandling: @@ -464,10 +460,9 @@ def test_skips_when_no_nixl_manager(self): @patch("modelexpress.load_strategy.base.publish_metadata_and_ready", side_effect=RuntimeError("gRPC fail")) def test_publish_failure_does_not_raise(self, _mock): from modelexpress.load_strategy.base import publish_metadata - ctx = _make_load_context() + ctx = _make_load_context(metadata_server_url="localhost:8001") ctx.nixl_manager = MagicMock() - with patch.dict(os.environ, {"MX_SERVER_ADDRESS": "localhost:8001"}): - publish_metadata(ctx) + publish_metadata(ctx) def test_unpublish_uses_worker_rank_for_heartbeat_lifecycle(self): from modelexpress.load_strategy.base import unpublish_metadata @@ -716,22 +711,14 @@ def _make_strategy(self): def test_unavailable_when_no_server_address(self, _mock): strategy = self._make_strategy() ctx = _make_load_context() - with patch.dict("os.environ", {}, clear=True): - assert strategy.is_available(ctx) is False - - @patch("modelexpress.load_strategy.rdma_strategy.is_nixl_available", return_value=True) - def test_available_when_mx_server_address_set(self, _mock): - strategy = self._make_strategy() - ctx = _make_load_context() - with patch.dict("os.environ", {"MX_SERVER_ADDRESS": "server:8001"}): - assert strategy.is_available(ctx) is True + assert strategy.is_available(ctx) is False @patch("modelexpress.load_strategy.rdma_strategy.is_nixl_available", return_value=True) - def test_available_when_model_express_url_set(self, _mock): + def test_available_when_context_metadata_server_url_set(self, _mock): strategy = self._make_strategy() ctx = _make_load_context() - with patch.dict("os.environ", {"MODEL_EXPRESS_URL": "server:8001"}): - assert strategy.is_available(ctx) is True + ctx.metadata_server_url = "server:8001" + assert strategy.is_available(ctx) is True @patch("modelexpress.load_strategy.rdma_strategy.is_nixl_available", return_value=False) def test_unavailable_when_nixl_not_available(self, _mock): @@ -749,8 +736,7 @@ def test_available_when_decentralized_backend_and_no_server_address(self, _mock) strategy = self._make_strategy() ctx = _make_load_context() ctx.mx_client.REQUIRES_P2P_METADATA = True - with patch.dict("os.environ", {}, clear=True): - assert strategy.is_available(ctx) is True + assert strategy.is_available(ctx) is True # --------------------------------------------------------------------------- @@ -776,6 +762,21 @@ def test_returns_empty_when_no_instances(self): result = strategy._find_source_instances(ctx) assert result == [] + def test_source_query_timeout_polls_until_matching_instance(self): + strategy = self._make_strategy() + ctx = _make_load_context() + ctx.load_config.source_query_timeout_s = 1 + inst = _make_instance_ref(worker_rank=ctx.worker_rank) + ctx.mx_client.list_sources.side_effect = [ + p2p_pb2.ListSourcesResponse(instances=[]), + p2p_pb2.ListSourcesResponse(instances=[inst]), + ] + with patch("modelexpress.load_strategy.rdma_strategy.time.sleep"), \ + patch("modelexpress.load_strategy.rdma_strategy.random.shuffle"): + result = strategy._find_source_instances(ctx) + assert result == [inst] + assert ctx.mx_client.list_sources.call_count == 2 + def test_returns_empty_when_list_sources_raises(self): strategy = self._make_strategy() ctx = _make_load_context() diff --git a/trtllm_patches/v1.3.0rc5/README.md b/trtllm_patches/v1.3.0rc5/README.md deleted file mode 100644 index 468baea07..000000000 --- a/trtllm_patches/v1.3.0rc5/README.md +++ /dev/null @@ -1,27 +0,0 @@ -# TRT-LLM 1.3.0rc5 PRESHARDED Patches - -Patches for TRT-LLM 1.3.0rc5 (used in Dynamo v1.0.0, image `karenc:dynamo-trtllm-v1.0.0-a9b6f95`). - -## File paths (inside container) - -```text -/opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/llmapi/llm_args.py -/opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/_torch/pyexecutor/model_loader.py -/opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/_torch/modules/linear.py -``` - -## Changes - -### llm_args.py (line ~2818) -Add `PRESHARDED = 3` to `LoadFormat` enum. - -### model_loader.py (line ~337, after DUMMY branch) -Add `elif load_format == LoadFormat.PRESHARDED:` branch that: -- Sets `_weights_presharded = True` on all Linear modules -- Calls `checkpoint_loader.load_weights(checkpoint_dir, mapping=self.mapping, model=model)` -- If result is empty dict, skips `model.load_weights()` - -### linear.py (functions at lines ~175, ~211, ~264) -In `load_weights_vanilla_helper`, `load_weights_fused_qkv_helper`, and -`load_weights_fused_gate_up_helper`: override `module.tp_size` to 1 when -`getattr(module, '_weights_presharded', False)` is True. diff --git a/trtllm_patches/v1.3.0rc5/apply_patches.py b/trtllm_patches/v1.3.0rc5/apply_patches.py deleted file mode 100644 index 0b6964bb8..000000000 --- a/trtllm_patches/v1.3.0rc5/apply_patches.py +++ /dev/null @@ -1,161 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""Apply PRESHARDED patches to TRT-LLM 1.3.0rc5. - -Run inside container: - python3 /tmp/apply_patches.py - -Patches 3 files to add LoadFormat.PRESHARDED support for ModelExpress P2P. -""" -import importlib.util -import re -import sys -from pathlib import Path - -spec = importlib.util.find_spec("tensorrt_llm") -SITE = Path(spec.submodule_search_locations[0]) if spec else Path("/opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm") - -def patch_llm_args(): - """Add PRESHARDED = 3 to LoadFormat enum.""" - p = SITE / "llmapi" / "llm_args.py" - text = p.read_text() - if "PRESHARDED" in text: - print("llm_args.py: already patched") - return - text = text.replace( - " VISION_ONLY = 2", - " VISION_ONLY = 2\n" - " # Weights already sharded per-rank (e.g. via ModelExpress P2P RDMA)\n" - " PRESHARDED = 3", - ) - p.write_text(text) - print("llm_args.py: patched (added PRESHARDED = 3)") - - -def patch_model_loader(): - """Add PRESHARDED branch to ModelLoader.load().""" - p = SITE / "_torch" / "pyexecutor" / "model_loader.py" - text = p.read_text() - if "PRESHARDED" in text: - print("model_loader.py: already patched") - return - - old = '''\ - elif load_format == LoadFormat.VISION_ONLY: - # Vision weights are already loaded within the model. - logger.info( - "LoadFormat.VISION_ONLY: skipping weight loading; using preloaded vision weights." - ) - - else:''' - - new = '''\ - elif load_format == LoadFormat.PRESHARDED: - for module in model.modules(): - if hasattr(module, 'tp_size'): - module._weights_presharded = True - weights = checkpoint_loader.load_weights( - checkpoint_dir, mapping=self.mapping, model=model) - if weights: - self.weight_mapper = checkpoint_loader.get_initialized_weight_mapper( - model, config) - self._call_load_weights(model.load_weights, weights, - self.weight_mapper) - else: - logger.info("PRESHARDED: weights injected directly, skipping load_weights()") - - elif load_format == LoadFormat.VISION_ONLY: - # Vision weights are already loaded within the model. - logger.info( - "LoadFormat.VISION_ONLY: skipping weight loading; using preloaded vision weights." - ) - - else:''' - - if old not in text: - print("model_loader.py: ERROR — cannot find insertion point", file=sys.stderr) - sys.exit(1) - text = text.replace(old, new) - p.write_text(text) - print("model_loader.py: patched (added PRESHARDED branch)") - - -def patch_linear(): - """Override tp_size to 1 when _weights_presharded is True in 3 helper functions.""" - p = SITE / "_torch" / "modules" / "linear.py" - text = p.read_text() - if "_weights_presharded" in text: - print("linear.py: already patched") - return - - helpers = [ - ("load_weights_vanilla_helper", "assert len(weights) == 1"), - ("load_weights_fused_qkv_helper", "if not allow_partial_loading:"), - ("load_weights_fused_gate_up_helper", "if not allow_partial_loading:"), - ] - - override_line = " _tp_size = 1 if getattr(module, '_weights_presharded', False) else module.tp_size\n" - count = 0 - - for func_name, anchor in helpers: - if func_name not in text: - print(f"linear.py: WARNING — cannot find {func_name}", file=sys.stderr) - continue - func_idx = text.find(func_name) - anchor_idx = text.find(anchor, func_idx) - if anchor_idx == -1: - print(f"linear.py: WARNING — cannot find anchor '{anchor}' in {func_name}", file=sys.stderr) - continue - line_start = text.rfind("\n", 0, anchor_idx) + 1 - text = text[:line_start] + override_line + text[line_start:] - count += 1 - - text = re.sub( - r'(load_weight_shard\([^)]*?,\s*)module\.tp_size', - r'\1_tp_size', - text, - ) - - p.write_text(text) - print(f"linear.py: patched ({count} helpers, replaced module.tp_size with _tp_size)") - - -def patch_worker_main(): - """Add publish_from_worker() call to worker_main() after worker creation.""" - p = SITE / "executor" / "worker.py" - text = p.read_text() - if "publish_from_worker" in text: - print("worker.py: already patched") - return - - anchor = '''\ - # Optionally disable GC (default: not disabled) - if os.getenv("TRTLLM_WORKER_DISABLE_GC", "0") == "1":''' - - patch = '''\ - # ModelExpress source: publish this rank's model params via NIXL - if os.environ.get("MODEL_EXPRESS_SOURCE"): - try: - from modelexpress.trtllm_live_transfer import publish_from_worker - publish_from_worker(worker) - except Exception as e: - logger.warning("ModelExpress publish_from_worker failed on rank %d: %s", mpi_rank(), e) - - # Optionally disable GC (default: not disabled) - if os.getenv("TRTLLM_WORKER_DISABLE_GC", "0") == "1":''' - - if anchor not in text: - print("worker.py: ERROR — cannot find insertion point", file=sys.stderr) - sys.exit(1) - text = text.replace(anchor, patch) - p.write_text(text) - print("worker.py: patched (added publish_from_worker hook)") - - -if __name__ == "__main__": - patch_llm_args() - patch_model_loader() - patch_linear() - patch_worker_main() - print("All patches applied successfully.") diff --git a/trtllm_patches/v1.3.0rc5/patch_model_loader.py b/trtllm_patches/v1.3.0rc5/patch_model_loader.py deleted file mode 100644 index 7894be404..000000000 --- a/trtllm_patches/v1.3.0rc5/patch_model_loader.py +++ /dev/null @@ -1,102 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""Patch model_loader.py for PRESHARDED P2P — Option 3 architecture. - -Source publishes weights BEFORE post_load_weights (pre-processed state). -Target receives via P2P and runs full post_load_weights normally. -No _mx_p2p_weights_loaded flag needed — both source and target run -the same post_load_weights transforms. - -apply_patches.py already adds the PRESHARDED block with load_weights skip. -This patch adds: -1. ModelExpress source publish hook BEFORE post_load_weights loop -2. Updates worker.py publish hook to skip if already published from model_loader -""" -import os -import sys - -target = "/opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/_torch/pyexecutor/model_loader.py" - -with open(target) as f: - content = f.read() - -# Patch 1: Insert source publish hook before post_load_weights loop -old1 = """ for module in model.modules(): - if hasattr(module, 'post_load_weights') and not getattr( - module, '_weights_removed', False): - module.post_load_weights()""" - -new1 = """ # ModelExpress source: publish pre-processed weights BEFORE - # post_load_weights so targets receive raw loaded state and can - # run their own post_load_weights() transforms. - if os.environ.get("MODEL_EXPRESS_SOURCE"): - try: - from modelexpress.trtllm_live_transfer import publish_model_params - publish_model_params(model) - model._mx_source_published = True - except Exception as e: - import logging - logging.getLogger("modelexpress").warning("ModelExpress publish failed: %s", e) - - for module in model.modules(): - if hasattr(module, 'post_load_weights') and not getattr( - module, '_weights_removed', False): - module.post_load_weights()""" - -if old1 in content and "publish_model_params" not in content: - content = content.replace(old1, new1) - print("patch_model_loader: patch 1 (source publish hook) applied") -elif "publish_model_params" in content: - print("patch_model_loader: patch 1 already applied") -else: - print("patch_model_loader: WARNING — patch 1 target not found", file=sys.stderr) - -# Ensure 'import os' exists at top of file -if "import os" not in content.split("\n")[0:20]: - content = "import os\n" + content - print("patch_model_loader: added 'import os'") - -with open(target, "w") as f: - f.write(content) -print("patch_model_loader: done") - - -# Patch 2: Update worker.py publish hook to skip if already published -worker_target = "/opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/executor/worker.py" -if os.path.exists(worker_target): - with open(worker_target) as f: - worker_content = f.read() - - old_worker = """ # ModelExpress source: publish this rank's model params via NIXL - if os.environ.get("MODEL_EXPRESS_SOURCE"): - try: - from modelexpress.trtllm_live_transfer import publish_from_worker - publish_from_worker(worker) - except Exception as e: - logger.warning("ModelExpress publish_from_worker failed on rank %d: %s", mpi_rank(), e)""" - - new_worker = """ # ModelExpress source: publish this rank's model params via NIXL. - # Skip if already published from ModelLoader.load() (pre-post_load_weights). - if os.environ.get("MODEL_EXPRESS_SOURCE"): - model = getattr(getattr(getattr(worker, 'engine', None), 'model_engine', None), 'model', None) - if model and getattr(model, '_mx_source_published', False): - logger.info("ModelExpress: already published from model_loader, skipping worker publish") - else: - try: - from modelexpress.trtllm_live_transfer import publish_from_worker - publish_from_worker(worker) - except Exception as e: - logger.warning("ModelExpress publish_from_worker failed on rank %d: %s", mpi_rank(), e)""" - - if old_worker in worker_content: - worker_content = worker_content.replace(old_worker, new_worker) - with open(worker_target, "w") as f: - f.write(worker_content) - print("patch_model_loader: worker.py patch applied (skip duplicate publish)") - elif "_mx_source_published" in worker_content: - print("patch_model_loader: worker.py already patched") - else: - print("patch_model_loader: WARNING — worker.py patch target not found", file=sys.stderr) -else: - print("patch_model_loader: worker.py not found at expected path") diff --git a/trtllm_patches/v1.3.0rc5/patch_tp_allgather.py b/trtllm_patches/v1.3.0rc5/patch_tp_allgather.py deleted file mode 100644 index bba233411..000000000 --- a/trtllm_patches/v1.3.0rc5/patch_tp_allgather.py +++ /dev/null @@ -1,154 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -""" -Patch TRT-LLM communicator.py to use chunked safe_allgather. - -Based on TRT-LLM PR #12174 by chienchunhung: - https://github.com/NVIDIA/TensorRT-LLM/pull/12174 - -Replaces raw comm.allgather(obj) in tp_allgather with safe_allgather() -that chunks transfers via MPI.Allgatherv to avoid 32-bit int overflow -and MPI_ERR_TRUNCATE with ob1 TCP BTL on GB200. -""" - -import os -import sys - -COMM_PATH = "/opt/dynamo/venv/lib/python3.12/site-packages/tensorrt_llm/_torch/distributed/communicator.py" - -SAFE_ALLGATHER = ''' -import math as _sa_math - -def safe_allgather(comm, obj, chunk_size: int = 4 * 1024 * 1024): - """Safely allgather potentially large objects by splitting into - fixed-size chunks using raw-byte MPI.Allgatherv. - - Based on TRT-LLM PR #12174. Avoids mpi4py 32-bit int overflow - in counts/displacements and pickle5 out-of-band buffer issues. - """ - if not ENABLE_MULTI_DEVICE: - return [obj] - if ENABLE_MULTI_DEVICE and MPI is None: - raise RuntimeError("mpi4py is required when ENABLE_MULTI_DEVICE is True") - if chunk_size <= 0: - raise ValueError("chunk_size must be > 0") - - rank = comm.Get_rank() - size = comm.Get_size() - - max_safe_chunk = np.iinfo(np.int32).max // size if size > 0 else chunk_size - if chunk_size > max_safe_chunk: - chunk_size = max_safe_chunk - - try: - payload = pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL) - my_n = len(payload) - except Exception as e: - _ = comm.allgather(-1) - raise RuntimeError(f"Rank {rank} serialization failed: {e}") from e - - lengths = np.array(comm.allgather(my_n), dtype=np.int64) - if (lengths < 0).any(): - raise RuntimeError("Serialization failed on at least one rank") - - displs = np.zeros(size, dtype=np.int64) - if size > 1: - displs[1:] = np.cumsum(lengths[:-1]) - - sendbuf_full = np.frombuffer(payload, dtype=np.uint8, count=my_n) - recvbuf = np.empty(lengths.sum(), dtype=np.uint8) - - max_len = lengths.max() if size > 0 else 0 - num_rounds = _sa_math.ceil(max_len / chunk_size) if max_len > 0 else 0 - - for r in range(num_rounds): - round_offs = r * chunk_size - counts_this_round = np.minimum( - np.maximum(lengths - round_offs, 0), chunk_size - ).astype(np.int32) - sent_so_far = np.minimum(lengths, round_offs) - - round_recvbuf = np.empty(counts_this_round.sum(), dtype=np.uint8) - round_displs = np.zeros(size, dtype=np.int32) - if size > 1: - round_displs[1:] = np.cumsum(counts_this_round[:-1]) - - send_part = sendbuf_full[ - sent_so_far[rank]:sent_so_far[rank] + counts_this_round[rank] - ] - - comm.Allgatherv( - [send_part, MPI.BYTE], - [round_recvbuf, counts_this_round, round_displs, MPI.BYTE], - ) - - src_offset = 0 - for i in range(size): - n = counts_this_round[i] - if n > 0: - dst_start = displs[i] + sent_so_far[i] - recvbuf[dst_start:dst_start + n] = ( - round_recvbuf[src_offset:src_offset + n] - ) - src_offset += n - - out = [] - for i in range(size): - sz = lengths[i] - if sz == 0: - out.append(None) - continue - start = displs[i] - blob = recvbuf[start:start + sz].tobytes() - try: - out.append(pickle.loads(blob)) - except Exception as e: - raise RuntimeError(f"Deserialization failed for rank {i}: {e}") from e - - return out -''' - -REPLACEMENT = ( - "def tp_allgather(self, obj):\n return self.tp_comm.allgather(obj)", - "def tp_allgather(self, obj, chunk_size: int = 4 * 1024 * 1024):\n return safe_allgather(self.tp_comm, obj, chunk_size=chunk_size)", -) - - -def main(): - if not os.path.exists(COMM_PATH): - print(f"communicator.py not found at {COMM_PATH}") - sys.exit(1) - - with open(COMM_PATH) as f: - src = f.read() - - if "safe_allgather" in src: - print("communicator.py already has safe_allgather — skipping") - return - - old, new = REPLACEMENT - if old in src: - src = src.replace(old, new) - print("Patched tp_allgather to use safe_allgather") - else: - print("WARNING: tp_allgather pattern not found — TRT-LLM version may differ") - sys.exit(1) - - lines = src.split("\n") - insert_idx = 0 - for i, line in enumerate(lines): - if line.startswith("import ") or line.startswith("from "): - insert_idx = i + 1 - - lines.insert(insert_idx, SAFE_ALLGATHER) - src = "\n".join(lines) - - with open(COMM_PATH, "w") as f: - f.write(src) - - print("Applied safe_allgather from TRT-LLM PR #12174 (chunk_size=4MB)") - - -if __name__ == "__main__": - main()