diff --git a/docker/Dockerfile b/docker/Dockerfile index 28efa072aaa..2bbdd67f4a6 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -74,6 +74,7 @@ apt-get install -y --no-install-recommends \ vim \ ccache \ protobuf-compiler \ + libzmq3-dev \ # Nsight apt install -y --no-install-recommends gnupg @@ -122,6 +123,69 @@ ENV PATH="/root/.local/bin:$PATH" RUN curl -LsSf https://astral.sh/uv/${UV_VERSION}/install.sh | sh && \ uv python install ${PYTHON_VERSION} +# KV cache transceiver (TRT-LLM prefill/decode disaggregation). +# +# TensorRT-LLM compiles every non-MPI cache-transceiver backend out of the wheel +# unless it is built against UCX, and the NIXL backend additionally needs NIXL +# headers/libs at cmake time (cpp/CMakeLists.txt only runs find_package(NIXL) +# inside its ENABLE_UCX block). backend=DEFAULT resolves to NIXL, so an image +# without NIXL can only serve an explicit backend=UCX. +# +# UCX itself ships in the NGC base image under /usr/local/ucx, but cmake's +# find_package(ucx) does not search that prefix on its own -- hence ucx_ROOT. +# libzmq (installed above) is the UCX wrapper's control plane: a hard cmake +# requirement to build, and a runtime dependency on every node that runs an +# engine, since the wrapper is dlopen()ed on the first KV transfer. +ENV ucx_ROOT=/usr/local/ucx +ENV LD_LIBRARY_PATH="/usr/local/ucx/lib:${LD_LIBRARY_PATH}" + +ARG NIXL_VERSION=v1.3.1 +RUN <<"EOF" bash -exu -o pipefail +if [[ ! -d /usr/local/ucx/lib/cmake/ucx ]]; then + echo "UCX missing from the base image; NIXL's UCX plugin cannot build" >&2 + exit 1 +fi + +# NIXL builds with meson/ninja and links Python bindings against pybind11. Like +# the other from-source builds in this repo (tools/build-custom-flashinfer.sh, +# tools/build-custom-vllm.sh), the build-time Python deps go into a uv venv +# rather than the system interpreter -- and this one is thrown away in the same +# layer, so nothing survives into the runtime environment. +uv venv /tmp/nixl-build-venv +uv pip install --python /tmp/nixl-build-venv/bin/python meson ninja pybind11 setuptools +export PATH="/tmp/nixl-build-venv/bin:${PATH}" + +GDS_PATH=/usr/local/cuda/targets/x86_64-linux +if [[ "$(uname -m)" != "x86_64" ]]; then + GDS_PATH=/usr/local/cuda/targets/sbsa-linux +fi +# NIXL links against the CUDA driver stub, which does not live on the default +# loader path inside the build container. +CUDA_SO_DIR=$(dirname "$(find /usr/local -name libcuda.so.1 | head -n1)") +export LD_LIBRARY_PATH="${LD_LIBRARY_PATH}:${CUDA_SO_DIR}" + +cd /tmp +git clone --depth 1 -b "${NIXL_VERSION}" https://github.com/ai-dynamo/nixl.git +cd nixl +meson setup builddir \ + -Ducx_path=/usr/local/ucx \ + -Dcudapath_lib=/usr/local/cuda/lib64 \ + -Dcudapath_inc=/usr/local/cuda/include \ + -Dgds_path="${GDS_PATH}" \ + -Dinstall_headers=true \ + -Ddisable_plugins=POSIX \ + -Dbuild_tests=false \ + -Dbuild_examples=false \ + --buildtype=release +cd builddir && ninja install +cd /tmp && rm -rf /tmp/nixl /tmp/nixl-build-venv + +test -f /opt/nvidia/nvda_nixl/include/nixl.h +EOF +# Both arch dirs are listed so this stays correct on x86_64 and aarch64; the +# missing one is simply skipped by the loader. +ENV LD_LIBRARY_PATH="/opt/nvidia/nvda_nixl/lib/x86_64-linux-gnu:/opt/nvidia/nvda_nixl/lib/aarch64-linux-gnu:/opt/nvidia/nvda_nixl/lib64:${LD_LIBRARY_PATH}" + # Disable usage stats by default for users who are sensitive to sharing usage. # Users are encouraged to enable if the wish. ENV RAY_USAGE_STATS_ENABLED=0 diff --git a/docs/design-docs/pd-disaggregation-rollout.md b/docs/design-docs/pd-disaggregation-rollout.md new file mode 100644 index 00000000000..0b8f0d142fa --- /dev/null +++ b/docs/design-docs/pd-disaggregation-rollout.md @@ -0,0 +1,233 @@ +# NemoRL<>TRTLLM Disaggregation Rollout (Experimental) + +> **Experimental**: PD disaggregation is wired for the TensorRT-LLM backend on the +> NeMo-Gym rollout path only. + +NeMo RL owns the *replica* (which engines exist, how they are placed, and their +lifecycle) and the *refit safety* (what happens when weights change mid-request). It +delegates both the *KV handshake* and the *routing inside a replica* to the inference +backend's own disaggregation front-end — for TensorRT-LLM, `OpenAIDisaggServer`. + +## Design + + + +### Topology + +A **replica** is a self-contained disaggregated fleet — M context engines plus K +generation engines — fronted by F `OpenAIDisaggServer` instances +(`num_frontend_workers`, default 1). N replicas are created, and each exposes F URLs: + +``` + replica 0 OpenAIDisaggServer fe_0 (one URL) --. + OpenAIDisaggServer fe_{F-1} (URL) --+ +NeMo-Gym --session affinity--> |-- ctx_0 .. ctx_{M-1} + `-- gen_0 .. gen_{K-1} + + replica 1 OpenAIDisaggServer fe_0 (one URL) --. + ... `-- ... +``` + +Every frontend of a replica is handed the *same* two address pools, so F changes nothing +about the engine layout — it only widens the relay in front of it. The frontend is a +single-process ASGI relay that terminates HTTP, strips the Gym-only request fields, +re-serializes the body, asks the routers, and makes two engine calls per turn. All of that +is CPU on one event loop under one GIL, so with long multi-turn prompts a single frontend +becomes the replica's turn-throughput ceiling well before its GPUs saturate. F > 1 spreads +that work over F processes. + +`TrtllmGeneration.dp_openai_server_base_urls` reports all N×F frontend URLs as one flat +list, so NeMo-Gym sees "N×F instances" exactly as it sees N without disaggregation. Its +existing per-session affinity (`responses_api_models/vllm_model/app.py::_resolve_client`, +a `session_id -> client` map seeded by `sha256(session_id) % len(clients)`) binds a +trajectory to one frontend — and therefore to one replica — for its whole lifetime. Since +each replica contributes exactly F of the N×F slots, replicas stay uniformly loaded. + +That stickiness is a **correctness** requirement, not just load-spreading. Each frontend +holds its own `ctx_router` state in-process, so two frontends of the same replica have no +shared view of which context engine holds a given conversation's prefix. If a trajectory's +later turn reached a different frontend, that frontend would route it to a context engine +that never saw the prefix — correct output, but a full re-prefill and the whole point of +`ctx_router: conversation` lost. + +Replicas are disjoint: no engine belongs to two of them. That is what lets every frontend +keep its routing state in-process and removes any need for cross-replica coordination — +`coordinator_url` is `None`, even with F > 1. + +**Everything below the replica boundary is opaque to NeMo RL.** Once a request reaches a +replica's `OpenAIDisaggServer`, that server alone decides which context engine serves the +prefill, which generation engine serves the decode, whether a context leg is needed at +all, and how the KV handshake between them is carried. NeMo RL never sees an individual +engine on the request path; it only composes each replica's two pools and selects the +router policies through `DisaggServerConfig` (see [Configuration](#configuration)). + +The cost of that boundary is that context affinity is no longer free. When a replica had +exactly one context engine, Gym's choice of URL *was* the choice of context engine. Now +`ctx_router` has to provide it, which is why it must be configured with a stateful policy +while `gen_router` need not be. + +Affinity is therefore established at three levels, and NeMo RL controls only the first: + +| Level | Decides | Driven by | State | +| --- | --- | --- | --- | +| Gym → frontend | which replica (and which of its F frontends) | `sha256(session_id)`, sticky | in Gym's `session_id -> client` map | +| frontend → context engine | which of the replica's M context engines | `ctx_router` | in that frontend's process | +| context engine → ADP rank | which attention-DP rank holds the prefix | `ConversationParams.conversation_id` relayed on the request | none — a function of the id | + +The third row is why `trtllm_http_server` forwards the conversation id (Gym's session id, +or the one the disagg service stamps onto `disaggregated_params` for its ctx/gen legs) to +`llm.generate_async`: under attention-DP each rank owns a separate KV pool, so without the +id a turn lands on an effectively random rank and re-prefills its history. + +### Forming a replica + +`TrtllmGeneration` lays engines out replica by replica, contexts before generations, so +membership is positional and needs no negotiation. Once every engine has initialised and +reported its address, it builds the address pools for each replica — `type='ctx'` entries +for that replica's context engines, `type='gen'` for its generation engines — and starts F +`DisaggServerActor`s against them, collecting one URL each. Those N×F URLs become +`dp_openai_server_base_urls`. + +Each frontend is a CPU-only Ray actor in its own process rather than a thread inside an +engine worker: it is the request hot path for the whole replica, and sharing a process +with an engine would couple the replica's routing latency to that one engine's load. Three +properties are worth noting, all of which exist so a crashed frontend can come back at the +*same* URL and Gym's sticky clients recover after their 5xx retries: + +- **Serving starts in `__init__`, not in `start()`.** Ray re-runs `__init__` on actor +restart but never replays method calls. `start()` only waits for `/health` and asserts the +serving thread is still alive. +- **The node pin is hard.** Frontends spread round-robin over their own replica's engine +nodes (`addrs[base + fe_idx % per_replica]`), so a frontend sits beside engines it talks +to, and restarts are unlimited. +- **Ports are deterministic**: `frontend_base_port + replica_idx * num_frontend_workers + +frontend_idx`. Both indices are needed — several replicas can share a node, and an offset +carrying only `frontend_idx` would have every replica's frontend 0 ask for the same port. +That failure is silent: `uvicorn`'s bind failure calls `sys.exit(1)`, which only unwinds +the daemon thread, and the `/health` probe is then answered 200 by whichever frontend won +the bind. Hence the liveness assert in `start()`. + +Every engine — context and generation alike — therefore has to run its own HTTP server, +because an address is the only way `OpenAIDisaggServer` can reach one. That is the sole +thing it is given about an engine. + +NeMo RL still created those engines and still holds their Ray actor handles, so refit does +not go through the disagg server at all: `TrtllmGeneration` drives the engines directly, as +it does without disaggregation. The same is true of sleep/wake, prefix-cache reset and +profiling. Delegation is confined to the request path; the control plane is unchanged. + +### Node affinity + +Placement matters at two scales, and both are preferences rather than hard requirements. + +**An engine's TP group should fit inside one node.** Tensor-parallel collectives are the +most bandwidth-hungry traffic in the system. When a per-role TP is wider than a node, the +engine should at least stay inside one *segment* — the `segment_size` nodes of a single +NVLink domain that `RayVirtualCluster` aligns placement to. + +**A replica should fit inside one node, or failing that one segment.** The ctx→gen KV +transfer happens entirely within a replica, so its cost is set by the slowest link any of +that replica's engines has to cross. A replica straddling a segment boundary pays network +bandwidth on every handoff, and it pays it per turn. + +Engines are laid out replica by replica with same-role engines contiguous, which keeps each role's GPUs adjacent. + +Which of the two boundaries actually costs depends on the platform: + +- **HGX, 8 GPUs per node** — the NVLink domain *is* the node, so crossing one drops +straight to the network. The node boundary is what matters. +- **GB200 NVL72** — Ray sees 4 GPUs per node and 18 of those nodes share one NVLink +domain, so crossing a node is cheap. The segment is the boundary worth respecting. + + + +## Refit safety + +Disaggregation does not introduce a new class of staleness. The aggregated path already +tolerates decoding KV that was built under previous weights: the only cache operation after +a weight update is `reset_prefix_cache()`, which clears the reusable prefix cache but not +the KV held by in-flight requests. + +`drain=True` is the one guarantee that does weaken. It blocks until `active_requests` and +`waiting_queue` are empty, and a context-only request whose KV has not been pulled yet is +in neither — once its forward finishes the engine releases its scheduler slot and parks it +in a separate `_requests_in_transfer` map, keeping only the KV blocks. So draining both +engines only empties their scheduler queues; a context engine can still be holding KV +blocks awaiting transfer, which a generation engine then decodes under the new weights. +This affects the synchronous path, where `drain=True` otherwise means "nothing holds KV +from the previous weights"; the in-flight path already accepts that. + +## Configuration + +```yaml +policy: + generation: + backend: trtllm + trtllm_cfg: + tensor_parallel_size: 2 + expose_http_server: true # every engine needs a disagg-capable endpoint + disaggregation: + enabled: true + num_context_engines: 2 # M — per replica + num_generation_engines: 3 # K — per replica, independent of M + cache_transceiver_backend: UCX # DEFAULT | UCX | NIXL | MOONCAKE | MPI + ctx_router: conversation # stateful: keeps a trajectory on one context engine + gen_router: load_balancing # stateless: placed locally, no coordinator + num_frontend_workers: 2 # F — disagg servers per replica; N*F <= 256 + frontend_base_port: 17300 # port = base + replica_idx * F + frontend_idx + # Per-role engine overrides, merged over trtllm_cfg. Any TRT-LLM kwarg + # goes here; TP and the MoE split are the ones that usually differ, + # and each role must satisfy moe_tp * moe_ep == its own TP. + ctx_trtllm_kwargs: + tensor_parallel_size: 4 + moe_tensor_parallel_size: 2 + moe_expert_parallel_size: 2 + gen_trtllm_kwargs: + tensor_parallel_size: 2 + moe_tensor_parallel_size: 1 + moe_expert_parallel_size: 2 + trtllm_kwargs: + kv_cache_config: + enable_block_reuse: true # see below +``` + +GPU budget: + +``` +replica width = M * context_tp + K * generation_tp +inference GPUs = num_replicas * replica width +``` + +`num_frontend_workers` does not appear here: frontends are CPU-only actors +(`num_cpus=1, num_gpus=0`), so F never changes the GPU budget or the engine layout. + +`TrtllmGeneration._plan_engines()` turns this into a per-engine `(role, width)` list — the +single source of truth everything else derives from: + +``` +[ctx_tp × M, gen_tp × K, ctx_tp × M, gen_tp × K, ...] + \______ replica 0 _____/ \______ replica 1 ... +``` + +Same-role engines are contiguous within a replica, which keeps each role's block of GPUs +adjacent. The replica *count* is derived from the inference cluster's size rather than +configured — `num_replicas = inference_GPUs / replica_width`, the same rule the DP-shard +count follows without disaggregation — so `replica_width` must divide the inference GPU +count. + +Each frontend's `DisaggServerConfig.node_id` must be set explicitly, and must be distinct +across *all* of them — not just across replicas. Its default is `uuid.getnode() % 256`, +documented as assuming one disagg server per machine, and we run N×F of them. It keys the +snowflake request-id mint (`process_id` is hardwired to 0 without a coordinator, and +`time.monotonic()` shares an origin across processes on one node), so a collision would +let two frontends mint the same `ctx_request_id`. We use +`replica_idx * num_frontend_workers + frontend_idx`; the field is 8 bits, hence the +asserted `num_replicas * num_frontend_workers <= 256`. + +## **Prerequisites** + +- **`ChatCompletionResponseChoice` needs a `token_ids` field.** Prompt token ids and +per-token logprobs already come back on the standard response, but generated token ids do +not — the chat choice carries only each token's decoded text. `CompletionResponseChoice` +already has the field; the chat one needs the same. + diff --git a/docs/index.md b/docs/index.md index b93e8e942fe..da2b6425b0c 100644 --- a/docs/index.md +++ b/docs/index.md @@ -373,6 +373,7 @@ design-docs/uv.md design-docs/dependency-management.md design-docs/chat-datasets.md design-docs/generation.md +design-docs/pd-disaggregation-rollout.md design-docs/dynamo-integration.md design-docs/sparse-delta-refit.md design-docs/checkpoint-engines.md diff --git a/examples/configs/grpo_math_1B_trtllm.yaml b/examples/configs/grpo_math_1B_trtllm.yaml index d377cbda4ed..cbb82101916 100644 --- a/examples/configs/grpo_math_1B_trtllm.yaml +++ b/examples/configs/grpo_math_1B_trtllm.yaml @@ -17,6 +17,68 @@ policy: # diverge from grpo.async_grpo (single source of truth). in_flight_weight_updates: ${grpo.async_grpo.in_flight_weight_updates} recompute_kv_cache_after_weight_updates: ${grpo.async_grpo.recompute_kv_cache_after_weight_updates} + # Prefill/decode disaggregation. A replica is num_context_engines context + # engines plus num_generation_engines generation engines, fronted by one + # disagg server that exposes the single URL NeMo-Gym talks to. The replica + # count is derived from the inference cluster size, not configured. + # Requires colocated.enabled=false and trtllm_cfg.expose_http_server=true. + disaggregation: + enabled: false + # Engines per replica; the two are independent, so the P:D ratio is free. + num_context_engines: 1 + num_generation_engines: 1 + # Stateful, so a trajectory's turns keep reaching the context engine + # holding its prefix: conversation | kv_cache_aware + ctx_router: conversation + # Stateless: a generation engine receives KV freshly on every turn, so + # it has nothing worth returning to: round_robin | load_balancing + gen_router: load_balancing + # TRT-LLM CacheTransceiverConfig backend for the KV handoff: + # DEFAULT | UCX | NIXL | MOONCAKE | MPI + cache_transceiver_backend: DEFAULT + # Frontend (disagg server) workers per replica. Each is its own actor + # with a distinct URL, so one frontend's CPU stops being the replica's + # turn-throughput ceiling. replicas * workers must be <= 256. + num_frontend_workers: 1 + # Base port for the frontends' deterministic ports + # (base + replica_idx * num_frontend_workers + frontend_idx). Keep it + # outside virtual_cluster's random master-port window (1400-1999). + frontend_base_port: 17300 + # Thin the ctx->gen relay: token ids as one base64 int32 string instead + # of a 30k-int JSON array, and drop the (redundant) message history. + gen_tokids_ctxbytes: false + gen_strip_message_history: false + # Render the chat template and tokenize on the frontends instead of the + # single ctx adapter process. Guarded by ctx-side shadow validation + # (NRL_TRTLLM_TOKENIZE_SHADOW_RATE). + frontend_tokenize: false + # KV transceiver tuning. null means "don't forward", so TRT-LLM keeps + # its own default for each. + # CPP | PYTHON. TRT-LLM's own default is "auto", which only adopts the + # model's preferred runtime when the effective backend supports it and + # silently falls back to C++ otherwise -- not what a hybrid Mamba model + # wants, since the recurrent-state handoff needs the Python (v2) one. + cache_transceiver_runtime: null + # MiB of bounce buffer, or 0 to keep the per-block path. Read only by + # the Python (v2) transceiver. + kv_cache_bounce_size_mb: null + max_tokens_in_buffer: null + # Milliseconds before an unfinished KV transfer is cancelled. TRT-LLM's + # default (60 s) is tuned for short prompts at low concurrency; large + # multi-turn workloads want a much larger value. + kv_transfer_timeout_ms: null + # Per-role overrides merged over trtllm_cfg (TP, the MoE split, and any + # other trtllm_cfg key), so the roles can have different shapes. A key + # set here WINS over trtllm_cfg -- pinning tensor_parallel_size below + # would make trtllm_cfg.tensor_parallel_size a no-op for the disagg + # layout, silently. Empty means both roles inherit trtllm_cfg unchanged. + # Each role must satisfy moe_tp * moe_ep == its own TP. For example: + # ctx_trtllm_kwargs: + # tensor_parallel_size: 4 + # gen_trtllm_kwargs: + # tensor_parallel_size: 2 + ctx_trtllm_kwargs: {} + gen_trtllm_kwargs: {} trtllm_kwargs: batch_wait_timeout_iters: 32 batch_wait_max_tokens_ratio: 0.5 diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index f31b01007f4..bbc35ac2496 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -1078,9 +1078,39 @@ def _spinup_nemo_gym(base_urls, model_name): ) elif generation_config["backend"] == "trtllm": trtllm_cfg = generation_config.get("trtllm_cfg", {}) - gpus_per_instance = trtllm_cfg[ - "tensor_parallel_size" - ] * trtllm_cfg.get("pipeline_parallel_size", 1) + disagg_cfg = trtllm_cfg.get("disaggregation") or {} + if disagg_cfg.get("enabled"): + # Under PD disaggregation the unit to keep inside one + # NVLink domain is the *replica*, not the engine: an + # engine's TP group all-reduces internally, but the KV + # cache handed from the replica's context engines to its + # generation engines crosses the transceiver on every + # turn. Sizing this by the engine (below) yields + # nodes_per_instance=1 whenever an engine fits in a node, + # which skips domain pinning entirely and lets a replica + # straddle racks -- correct, but with the KV transfer + # demoted from NVLink to InfiniBand. + # + # No pipeline_parallel_size factor: TrtllmGeneration + # asserts pp == 1, so folding it in would only suggest a + # dimension this backend does not have. + def _role_tp(role: str) -> int: + overrides = disagg_cfg.get(f"{role}_trtllm_kwargs") or {} + return int( + overrides.get( + "tensor_parallel_size", + trtllm_cfg["tensor_parallel_size"], + ) + ) + + gpus_per_instance = int( + disagg_cfg["num_context_engines"] * _role_tp("ctx") + + disagg_cfg["num_generation_engines"] * _role_tp("gen") + ) + else: + gpus_per_instance = trtllm_cfg[ + "tensor_parallel_size" + ] * trtllm_cfg.get("pipeline_parallel_size", 1) elif generation_config["backend"] == "dynamo": gpus_per_instance = DynamoConfig.model_validate( generation_config diff --git a/nemo_rl/distributed/virtual_cluster.py b/nemo_rl/distributed/virtual_cluster.py index b3c4f96af94..efb663f400b 100644 --- a/nemo_rl/distributed/virtual_cluster.py +++ b/nemo_rl/distributed/virtual_cluster.py @@ -1187,6 +1187,62 @@ def _get_sorted_bundle_indices(self) -> Optional[list[int]]: ) return reordered_bundle_indices + def get_topology_sorted_pg_indices(self) -> list[int]: + """Placement group indices ordered by physical topology. + + The unified-PG path sorts *bundles* (see ``_get_sorted_bundle_indices``), + but with per-node placement groups the caller consumes whole PGs, and Ray + decides which physical node backs each one. Two PGs that are adjacent in + creation order can therefore sit in different NVLink domains, which for a + consumer whose unit spans several nodes (a TRT-LLM disaggregated replica, + say) silently demotes its inter-node traffic from NVLink to InfiniBand. + + Sorting by ``(nvlink_domain's smallest topo_rank, topo_rank)`` puts nodes + of one domain next to each other, so consecutive PGs belong to the same + domain wherever the allocation allows it. Ordering *within* a domain is + not a bandwidth concern (an NVL72 fabric is uniform) but keeps a given + config reproducible across runs. + + Returns: + Indices into ``get_placement_groups()``. Falls back to the existing + order when the cluster is CPU-only, is a single unified PG, or has no + topology information -- ordering is an optimization, never a + requirement. + """ + pgs = self.get_placement_groups() + identity = list(range(len(pgs))) + if not self.use_gpus or len(pgs) <= 1: + return identity + + topology = get_ray_cluster_topology() + if not any(domain != NVLINK_DOMAIN_UNKNOWN for domain, _ in topology.values()): + return identity + + node_of_pg: list[str] = [] + for pg in pgs: + bundles_to_node = placement_group_table(pg)["bundles_to_node_id"] + if not bundles_to_node: + return identity + node_of_pg.append(next(iter(bundles_to_node.values()))) + if any(node_id not in topology for node_id in node_of_pg): + return identity + + # Rank domains by their smallest topo_rank so domain order is stable and + # follows the fabric, not dict iteration order. + domain_rank: dict[str, int] = {} + for domain, topo_rank in topology.values(): + if domain not in domain_rank or topo_rank < domain_rank[domain]: + domain_rank[domain] = topo_rank + + return sorted( + identity, + key=lambda i: ( + domain_rank[topology[node_of_pg[i]][0]], + topology[node_of_pg[i]][1], + i, + ), + ) + def shutdown(self) -> bool: """Cleans up and releases all resources associated with this virtual cluster. diff --git a/nemo_rl/distributed/worker_groups.py b/nemo_rl/distributed/worker_groups.py index 3cfd2d405ca..816cc9d47c5 100644 --- a/nemo_rl/distributed/worker_groups.py +++ b/nemo_rl/distributed/worker_groups.py @@ -354,6 +354,7 @@ def __init__( bundle_indices_list: Optional[list[tuple[int, list[int]]]] = None, sharding_annotations: Optional[NamedSharding] = None, env_vars: dict[str, str] = {}, + dp_leader_worker_indices: Optional[list[int]] = None, ): """Initialize a group of distributed Ray workers. @@ -367,6 +368,15 @@ def __init__( Each tuple defines a tied group of workers placed on the same node. If provided, workers_per_node is ignored. sharding_annotations: NamedSharding object representing mapping of named axes to ranks (i.e. for TP, PP, etc.) + dp_leader_worker_indices: Global rank of the worker that owns each data + parallel shard. Defaults to the leader of every tied + group, i.e. a tied group *is* a DP shard -- true + whenever one model owner serves a whole rollout shard. + Pass it when the caller's DP layout is coarser than + the tied groups, e.g. TRT-LLM PD disaggregation, where + a replica is several engines (context + generation) + and counting tied groups would overcount DP. Must be a + subset of the tied group leaders, in ascending order. """ self._workers: list[ray.actor.ActorHandle] = [] self._worker_metadata: list[dict[str, Any]] = [] @@ -433,6 +443,20 @@ def __init__( env_vars=env_vars, ) + if dp_leader_worker_indices is not None: + # A DP leader has to be a model owner -- it is the worker every + # shard-level call is routed to. Checking against the tied group + # leaders derived above turns a wrong layout into a startup error + # instead of requests landing on a worker that never built a model. + unknown = set(dp_leader_worker_indices) - set(self.dp_leader_worker_indices) + if unknown or dp_leader_worker_indices != sorted(dp_leader_worker_indices): + raise ValueError( + f"dp_leader_worker_indices {dp_leader_worker_indices} must be an " + f"ascending subset of the tied group leaders " + f"{self.dp_leader_worker_indices}." + ) + self.dp_leader_worker_indices = dp_leader_worker_indices + def get_dp_leader_worker_idx(self, dp_shard_idx: int) -> int: """Returns the index of the primary worker for a given data parallel shard.""" if not 0 <= dp_shard_idx < len(self.dp_leader_worker_indices): @@ -481,6 +505,10 @@ def _create_workers_from_bundle_indices( ) else: py_executable = actor_python_env + # Resolved interpreter for this group's workers, kept so sidecar actors + # that must import the same packages (e.g. the TRT-LLM disagg server) + # can join the venv instead of landing in the driver's environment. + self.py_executable = py_executable # Count total workers self.world_size = sum(len(indices) for _, indices in bundle_indices_list) diff --git a/nemo_rl/models/generation/trtllm/config.py b/nemo_rl/models/generation/trtllm/config.py index 10d3b16d205..f4f8658935e 100644 --- a/nemo_rl/models/generation/trtllm/config.py +++ b/nemo_rl/models/generation/trtllm/config.py @@ -12,11 +12,119 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Any, NotRequired, TypedDict +from typing import Any, Literal, NotRequired, Optional, TypedDict + +from pydantic import BaseModel, Field from nemo_rl.models.generation.interfaces import GenerationConfig +class TrtllmDisaggConfig(BaseModel, extra="allow"): + """Prefill/decode disaggregation. + + A *replica* is ``num_context_engines`` context engines plus + ``num_generation_engines`` generation engines, fronted by one + ``OpenAIDisaggServer`` that exposes the single URL NeMo-Gym talks to. The + replica *count* is not configured: it follows from the inference cluster's + size, the same way the DP-shard count does without disaggregation. + + Requires non-colocated generation: colocated sleeps the engines between + rollouts, and a replica's context and generation engines must be resident + together for the KV transceiver to work. + """ + + enabled: bool = False + + # Engines per replica. The two are independent, so the P:D ratio is free. + num_context_engines: int = 1 + num_generation_engines: int = 1 + + # Routing inside a replica, decided entirely by the disagg server. + # + # The context router must be *stateful* so a trajectory's turns return to + # the engine holding its prefix -- that engine accumulates the prefix across + # turns and only prefills the delta, so sending a later turn elsewhere + # throws the work away. + # + # The generation router need not be: a generation engine receives KV freshly + # from the context engine on every turn, so it has nothing worth returning + # to, and a wrong load guess only costs transient skew. Keeping it stateless + # also keeps placement local, with no coordinator process. + # + # Both are ``Literal`` rather than ``str`` because a plausible-but-wrong + # value is the dangerous case: ``ctx_router="round_robin"`` parses fine and + # silently throws away the prefix affinity the context engines depend on. + ctx_router: Literal["conversation", "kv_cache_aware"] = "conversation" + gen_router: Literal["round_robin", "load_balancing"] = "load_balancing" + + # Frontend (disagg server) workers per replica. Each is its own + # DisaggServerActor with a distinct URL; NeMo-Gym's per-session client + # selection shards conversations across them, so one frontend's CPU stops + # being the replica's turn-throughput ceiling. 1 = single-frontend + # behavior. replicas * workers must be <= 256 (snowflake node_id space). + num_frontend_workers: int = 1 + # Relay ctx->gen prompt token ids as one base64 int32 string instead of a + # 30k-int JSON array (TRT-LLM DisaggServerConfig.gen_tokids_ctxbytes). + gen_tokids_ctxbytes: bool = False + # Strip the conversation history from the generation leg; the relayed + # token ids carry the full prefix, so the generation adapter never needs + # the messages (DisaggServerConfig.gen_strip_message_history). + gen_strip_message_history: bool = False + # Frontends render the chat template and tokenize (via the adapters' + # exact shared pipeline) and attach prompt_token_ids_b64 to the ctx leg, + # so the single ctx adapter process does no template work. Guarded by + # ctx-side shadow validation (NRL_TRTLLM_TOKENIZE_SHADOW_RATE). + frontend_tokenize: bool = False + # Base port for the frontend workers' deterministic ports + # (base + replica_idx * num_frontend_workers + frontend_idx). Deterministic + # so a restarted frontend actor re-binds the SAME port and its URL stays + # valid; keep the range outside virtual_cluster's random master-port window + # (1400-1999). + frontend_base_port: int = 17300 + + # Mapped onto TRT-LLM's CacheTransceiverConfig. + cache_transceiver_backend: Literal["DEFAULT", "UCX", "NIXL", "MOONCAKE", "MPI"] = ( + "DEFAULT" + ) + + # "CPP" | "PYTHON". TRT-LLM defaults to "auto", which only adopts + # the model's preferred runtime when the effective backend supports it and + # silently falls back to the C++ transceiver otherwise -- and that fallback + # is not what a hybrid Mamba model wants: the recurrent-state handoff needs + # the Python (v2) transceiver. Left unset here so TRT-LLM keeps its own + # default; set it explicitly to force one. + cache_transceiver_runtime: Optional[Literal["CPP", "PYTHON"]] = None + + # MiB of bounce buffer, or 0 to keep the per-block path. Bounce coalesces a + # request's scattered per-block KV into one contiguous fabric-VMM buffer and + # issues a single multi-rail NIXL write, which sidesteps registering every + # VMM-split block descriptor individually -- the step that fails here with + # "registerMem: registration failed for the specified or all potential + # backends". Only the Python (v2) transceiver reads it. + kv_cache_bounce_size_mb: Optional[int] = None + max_tokens_in_buffer: Optional[int] = None + + # Milliseconds before an unfinished KV transfer is cancelled on either + # side. TRT-LLM's default (60 s) is tuned for short prompts at low + # concurrency; at high rollout concurrency the ctx-side timeout can fire + # in bulk and the resulting cancel/retry churn stresses the transceiver, + # so large multi-turn workloads want a much larger value. + kv_transfer_timeout_ms: Optional[int] = None + + # Per-role overrides merged over trtllm_cfg. Any trtllm_cfg key goes here -- + # tensor_parallel_size and the MoE split are the ones that usually differ, + # and each role must satisfy moe_tp * moe_ep == its own TP. A role may also + # carry its own ``trtllm_kwargs`` (including ``kv_cache_config``) when + # prefill and decode want different engine tuning. Genuinely an arbitrary + # TRT-LLM passthrough, so ``dict[str, Any]`` is the right type here. + ctx_trtllm_kwargs: dict[str, Any] = Field(default_factory=dict) + gen_trtllm_kwargs: dict[str, Any] = Field(default_factory=dict) + + def role_trtllm_kwargs(self, role: Literal["ctx", "gen"]) -> dict[str, Any]: + """One role's engine overrides, selected without attribute reflection.""" + return self.ctx_trtllm_kwargs if role == "ctx" else self.gen_trtllm_kwargs + + class TrtllmSpecificArgs(TypedDict): tensor_parallel_size: int model_name: NotRequired[str] @@ -43,6 +151,7 @@ class TrtllmSpecificArgs(TypedDict): # grpo.async_grpo so they cannot diverge). in_flight_weight_updates: NotRequired[bool] recompute_kv_cache_after_weight_updates: NotRequired[bool] + disaggregation: NotRequired[TrtllmDisaggConfig] default_chat_template_kwargs: NotRequired[dict[str, Any]] # TRT-LLM's registered parser names: # "qwen3" -> Qwen3ToolParser (JSON format: {"name":..., "arguments":{...}}) @@ -57,3 +166,18 @@ class TrtllmConfig(GenerationConfig): # covered by TrtllmSpecificArgs (e.g. sampler_type, enable_attention_dp). # Spread into the engine constructor as `**trtllm_kwargs`. trtllm_kwargs: NotRequired[dict[str, Any]] + + +def resolve_trtllm_disagg_config(config: TrtllmConfig) -> TrtllmDisaggConfig: + """Validate ``trtllm_cfg.disaggregation`` into its schema. + + Absent is the same as present-and-disabled: every field carries its default + on :class:`TrtllmDisaggConfig`, so callers read attributes unconditionally + instead of re-deriving a default per key at the call site. + """ + raw = config["trtllm_cfg"].get("disaggregation") + if raw is None: + return TrtllmDisaggConfig() + if isinstance(raw, TrtllmDisaggConfig): + return raw + return TrtllmDisaggConfig.model_validate(raw) diff --git a/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py new file mode 100644 index 00000000000..3a091cf8e2f --- /dev/null +++ b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py @@ -0,0 +1,654 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""One replica's disaggregation front-end, backed by TRT-LLM's ``OpenAIDisaggServer``. + +Started the same way as :mod:`trtllm_http_server`: a uvicorn app in a daemon +thread, returning the URL NeMo-Gym will talk to. That server owns everything +below the replica boundary — which context engine runs the prefill, which +generation engine runs the decode, and the KV handshake between them. NeMo RL +only hands it the two address pools and the router policies. +""" + +import json +import logging +import threading +import time +from typing import Any, Optional + +import ray + +logger = logging.getLogger(__name__) + +__all__ = [ + "DisaggServerActor", + "DisaggServerActorImpl", + "attach_rollout_fields", + "build_config", + "generation_token_ids", + "start_server", + "wait_ready", +] + + +def build_config( + ctx_addrs: list[tuple[str, int]], + gen_addrs: list[tuple[str, int]], + *, + node_id: int, + ctx_router: str, + gen_router: str, + gen_tokids_ctxbytes: bool = False, + gen_strip_message_history: bool = False, +) -> Any: + """Assemble the ``DisaggServerConfig`` for one replica. + + Args: + ctx_addrs / gen_addrs: ``(hostname, port)`` of every engine in this + replica. An address is the only thing the disagg server is given + about an engine. + node_id: Must be distinct per replica. TRT-LLM's default is + ``uuid.getnode() % 256``, documented as assuming a single disagg + server per machine, and we run one per replica. + ctx_router: Stateful, so a trajectory's turns keep reaching the context + engine holding its prefix. + gen_router: Stateless, so placement stays local to this replica and no + coordinator process is needed. + """ + from tensorrt_llm.llmapi.disagg_utils import ( + CtxGenServerConfig, + DisaggServerConfig, + RouterConfig, + ) + + server_configs = [ + CtxGenServerConfig(type="ctx", hostname=host, port=port) + for host, port in ctx_addrs + ] + [ + CtxGenServerConfig(type="gen", hostname=host, port=port) + for host, port in gen_addrs + ] + + config = DisaggServerConfig( + server_configs=server_configs, + ctx_router_config=RouterConfig(type=ctx_router), + gen_router_config=RouterConfig(type=gen_router), + node_id=node_id, + ) + # Post-construction assignment, same as TRT-LLM's own YAML parser + # (llmapi/disagg_utils.py extract_disagg_cfg): thins the gen leg so the + # relay stops re-serializing a 30k-int id array and the full history. + config.gen_tokids_ctxbytes = gen_tokids_ctxbytes + config.gen_strip_message_history = gen_strip_message_history + return config + + +# Request fields clients send for vLLM that TRT-LLM's ChatCompletionRequest +# does not declare. Its models are extra="forbid", so leaving them in means a +# 422/400 on every request. `required_prefix_token_ids` is the vLLM server's +# prefix-override extension (used by the MLPerf warmup); the TRT-LLM engine +# adapter derives the on-policy prefix from the assistant messages instead and +# ignores the top-level field, so dropping it here keeps both paths identical. +_GYM_ONLY_REQUEST_FIELDS = ( + "return_tokens_as_token_ids", + "return_token_ids", + "required_prefix_token_ids", +) + +# Prefix vLLM uses when asked to report tokens as ids, which NeMo-Gym parses. +_TOKEN_ID_PREFIX = "token_id:" + +# Paths whose bodies are validated against TRT-LLM's extra="forbid" models. +_ADAPTED_PATHS = frozenset({"/v1/chat/completions", "/v1/completions"}) + + +def generation_token_ids(choice: dict[str, Any]) -> Optional[list[int]]: + """Generated token ids, from whichever declared field carries them. + + ``ChatCompletionResponseChoice.token_ids`` is the clean home, but it does + not exist upstream yet. Until it does they ride in + ``logprobs.content[].token`` using vLLM's ``token_id:N`` encoding, which is + a declared string field and therefore survives the hop. + """ + if choice.get("token_ids"): + return list(choice["token_ids"]) + + content = (choice.get("logprobs") or {}).get("content") or [] + ids = [] + for entry in content: + token = entry.get("token") or "" + if not token.startswith(_TOKEN_ID_PREFIX): + return None + ids.append(int(token[len(_TOKEN_ID_PREFIX) :])) + return ids or None + + +def attach_rollout_fields(payload: dict[str, Any]) -> dict[str, Any]: + """Re-attach, in place, the fields NeMo-Gym reads off ``choices[].message``. + + The counterpart of the aggregated path's ``msg_dict`` assignment in + :mod:`trtllm_http_server`: there the engine answers Gym directly and writes + the fields onto the message, here the engine had to route them through + declared response fields to survive the disagg server's ``extra="forbid"`` + re-validation, so they are moved back. The two must agree field for field -- + a message missing ``generation_token_ids`` is skipped outright by + ``nemo_gym.py``, silently dropping that turn's training data -- which is why + this lives at module level with no ``tensorrt_llm`` import in its way. + """ + choices = payload.get("choices") or [] + if not choices: + return payload + + choice = choices[0] + message = choice.get("message") + if not isinstance(message, dict): + return payload + + if payload.get("prompt_token_ids") is not None: + message["prompt_token_ids"] = payload["prompt_token_ids"] + + token_ids = generation_token_ids(choice) + if token_ids is not None: + message["generation_token_ids"] = token_ids + + content = (choice.get("logprobs") or {}).get("content") + if content: + message["generation_log_probs"] = [entry.get("logprob") for entry in content] + + return payload + + +class _DropGymOnlyRequestFields: + """ASGI middleware stripping the vLLM-only fields from request bodies. + + Deliberately raw ASGI rather than a Starlette ``BaseHTTPMiddleware``: that + class passes the downstream app its own captured receive channel, so + reassigning ``request._receive`` never reaches FastAPI's validation and the + request is rejected anyway. Replacing ``receive`` here is what the route + actually reads. + """ + + def __init__(self, app: Any, tokenize_fn: Any = None) -> None: + self.app = app + # Phase-2 frontend tokenize: async callable(payload) -> b64 ids or + # None. Runs here because this middleware already pays the + # loads/dumps round trip for every body. + self._tokenize_fn = tokenize_fn + + async def __call__(self, scope: Any, receive: Any, send: Any) -> None: + if scope.get("type") != "http" or scope.get("path") not in _ADAPTED_PATHS: + await self.app(scope, receive, send) + return + + chunks: list[bytes] = [] + while True: + message = await receive() + if message["type"] == "http.disconnect": + return + chunks.append(message.get("body", b"")) + if not message.get("more_body", False): + break + body = b"".join(chunks) + + try: + payload = json.loads(body) + except ValueError: + payload = None + + mutated = False + if isinstance(payload, dict) and any( + field in payload for field in _GYM_ONLY_REQUEST_FIELDS + ): + for field in _GYM_ONLY_REQUEST_FIELDS: + payload.pop(field, None) + mutated = True + if ( + self._tokenize_fn is not None + and isinstance(payload, dict) + and payload.get("messages") + and "prompt_token_ids_b64" not in payload + and "prompt_token_ids" not in payload + and payload.get("disaggregated_params") is None + ): + # Inbound gym request (the engine-facing legs carry + # disaggregated_params and never traverse this app). Attach the + # rendered+spliced prompt ids so the ctx adapter skips its + # chat-template work entirely; prompt_token_ids_b64 is a declared + # request field end-to-end. + try: + b64 = await self._tokenize_fn(payload) + except Exception: # noqa: BLE001 - fall back to adapter-side tokenize + logger.exception( + "frontend tokenize failed; falling back to adapter-side" + ) + b64 = None + if b64 is not None: + payload["prompt_token_ids_b64"] = b64 + mutated = True + if mutated: + body = json.dumps(payload).encode() + # Content-Length must follow the body when it changed; a stale value + # makes any proxy in front of this server truncate or hang. + headers = [ + (name, value) + for name, value in scope["headers"] + if name.lower() != b"content-length" or not mutated + ] + if mutated: + headers.append((b"content-length", str(len(body)).encode())) + scope = {**scope, "headers": headers} + + delivered = False + + async def _receive() -> dict[str, Any]: + nonlocal delivered + if delivered: + return {"type": "http.disconnect"} + delivered = True + return {"type": "http.request", "body": body, "more_body": False} + + await self.app(scope, _receive, send) + + +def _build_adaptor_class() -> type: + """Build ``OpenAIDisaggServerAdaptor`` lazily. + + Deferred so importing this module does not require tensorrt_llm. + """ + import aiohttp + from fastapi import HTTPException, Request, Response + from fastapi.responses import JSONResponse + from tensorrt_llm.serve.openai_disagg_server import OpenAIDisaggServer + from tensorrt_llm.serve.openai_protocol import UCompletionRequest + + class OpenAIDisaggServerAdaptor(OpenAIDisaggServer): + """``OpenAIDisaggServer`` speaking NeMo-Gym's dialect at both edges. + + The disagg server validates strictly in both directions -- requests into + ``ChatCompletionRequest``, engine responses into + ``ChatCompletionResponse``, both ``extra="forbid"``. NeMo-Gym, which was + written against vLLM, sends one request field TRT-LLM does not declare + and reads rollout fields that are not part of the OpenAI schema. This + subclass reconciles the two without touching either side: + + * inbound -- drop the vLLM-only request fields before FastAPI validates + * outbound -- re-attach the rollout fields Gym reads off the message + + Note the outbound translation only moves fields that already survived + the engine -> disagg-server hop. It cannot rescue a field the engine + emitted that ``ChatCompletionResponse`` does not declare; those have to + travel in a declared field (see ``_generation_token_ids``). + """ + + def __init__(self, *args: Any, **kwargs: Any) -> None: + self._tokenize_fn = kwargs.pop("tokenize_fn", None) + super().__init__(*args, **kwargs) + self._uvicorn: Any = None + self._install_request_adaptor() + + async def __call__(self, host: str, port: int, sockets: Any = None) -> None: + """Serve, keeping a handle on uvicorn so shutdown can stop it. + + The base class builds its ``uvicorn.Server`` as a local, so there is + nothing to signal ``should_exit`` on otherwise. + """ + import uvicorn + + config = uvicorn.Config( + self.app, host=host, port=port, log_level="info", timeout_keep_alive=10 + ) + self._uvicorn = uvicorn.Server(config) + await self._uvicorn.serve(sockets=sockets) + + def request_shutdown(self) -> None: + if self._uvicorn is not None: + self._uvicorn.should_exit = True + + # -------------------------------------------------------------- # + # Inbound + # -------------------------------------------------------------- # + + def _install_request_adaptor(self) -> None: + # Pure ASGI, not @app.middleware("http"): BaseHTTPMiddleware hands + # the downstream app the receive channel it captured itself, so + # rewriting request._receive there is invisible to FastAPI's + # validation and every request still fails with extra_forbidden. + # Wrapping receive at the ASGI layer is what actually replaces the + # body the route sees. + self.app.add_middleware( + _DropGymOnlyRequestFields, tokenize_fn=self._tokenize_fn + ) + + # -------------------------------------------------------------- # + # Errors + # -------------------------------------------------------------- # + + def _handle_exception(self, exception: BaseException) -> None: + """Pass a downstream 4xx through instead of masking it as a 500. + + The base implementation only re-raises ``HTTPException``; a 4xx + from a context/generation engine arrives as an + ``aiohttp.ClientResponseError`` and falls into the catch-all that + turns it into ``500 Internal server error``. The aggregated server + returns that same rejection to the caller as a 4xx, so without this + a deterministic client error -- most commonly the + ``context length exceeded`` guard in ``trtllm_http_server`` -- looks + like a server fault under disaggregation only, and Gym's masking + statistics stop being comparable between the two paths. + """ + if ( + isinstance(exception, aiohttp.ClientResponseError) + and 400 <= exception.status < 500 + ): + self._perf_metrics_collector.http_exceptions.inc() + raise HTTPException( + status_code=exception.status, detail=exception.message + ) + super()._handle_exception(exception) + + # -------------------------------------------------------------- # + # Outbound + # -------------------------------------------------------------- # + + def _wrap_entry_point( + self, entry_point: Any, request_type: type = UCompletionRequest + ) -> Any: + inner = super()._wrap_entry_point(entry_point, request_type) + + async def wrapper(req: request_type, raw_req: Request) -> Response: # type: ignore[valid-type] + response = await inner(req, raw_req) + if req.stream or not isinstance(response, JSONResponse): + return response + # The translation itself is module-level and tensorrt_llm-free + # so it can be unit-tested against the exact payload + # trtllm_http_server emits; only the JSONResponse re-wrap + # belongs to this subclass. + return JSONResponse( + content=attach_rollout_fields(json.loads(response.body)), + status_code=response.status_code, + ) + + return wrapper + + return OpenAIDisaggServerAdaptor + + +def _build_frontend_tokenizer( + model_name: str, server_template_kwargs: dict[str, Any] +) -> Any: + """Async ``payload -> b64 prompt ids`` using the adapters' exact pipeline. + + Loads the tokenizer and HF config the same way the engine adapters do + (``AutoTokenizer/AutoConfig.from_pretrained(model_name, + trust_remote_code=True)``) and renders through the shared + ``build_spliced_prompt_ids`` -- the ids must be bit-identical to what the + ctx adapter would compute, since a divergence silently trains on + off-policy tokens (the adapter shadow-validates a sample under + NRL_TRTLLM_TOKENIZE_SHADOW_RATE). + """ + import base64 + + import numpy as np + from tensorrt_llm.serve.openai_protocol import ChatCompletionRequest + from transformers import AutoConfig, AutoTokenizer + + from nemo_rl.models.generation.trtllm.trtllm_http_server import ( + build_spliced_prompt_ids, + ) + + tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True) + model_config = AutoConfig.from_pretrained(model_name, trust_remote_code=True) + + async def _tokenize(payload: dict[str, Any]) -> str: + # Normalize through the SAME pydantic model and dump flags the disagg + # client uses toward the engines (model_dump(mode="json", + # exclude_unset=True), openai_client.py). The ctx adapter renders the + # POST-VALIDATION dicts -- pydantic inserts defaults and fixes key + # order (e.g., tools entries become {"type": "function", "function": + # ...}), and the chat template serializes them verbatim, so rendering + # the raw gym payload instead diverges at the tools JSON on turn-1 + # prompts (caught by shadow validation). + norm = ChatCompletionRequest(**payload).model_dump( + mode="json", exclude_unset=True + ) + per_request = norm.get("chat_template_kwargs") or {} + effective = {**server_template_kwargs, **per_request} + ids = await build_spliced_prompt_ids( + norm.get("messages") or [], + norm.get("tools"), + tokenizer, + model_config, + effective, + ) + return base64.b64encode(np.asarray(ids, dtype=np.int32).tobytes()).decode() + + return _tokenize + + +def start_server( + config: Any, + host: str = "0.0.0.0", + port: int = 0, + req_timeout_secs: int = 1800, + tokenize_fn: Any = None, +) -> "tuple[threading.Thread, str, Any]": + """Start the disagg server in a daemon thread and return (thread, base_url, server).""" + import asyncio + + from nemo_rl.distributed.virtual_cluster import ( + _get_free_port_local, + _get_node_ip_local, + ) + + if port == 0: + port = _get_free_port_local() + + node_ip = _get_node_ip_local() + base_url = f"http://{node_ip}:{port}/v1" + + # OpenAIDisaggServer.register_routes() builds a prometheus + # MultiProcessCollector, which raises unless PROMETHEUS_MULTIPROC_DIR points + # at a real directory. TRT-LLM's own entrypoint calls this helper first + # (tensorrt_llm/commands/serve.py); we construct the server directly, so we + # have to. It keeps the TemporaryDirectory alive in a module global, which + # is also what stops it from being collected while the server runs. + from tensorrt_llm._utils import set_prometheus_multiproc_dir + + set_prometheus_multiproc_dir() + + # coordinator_url=None: this server owns its routing state in-process. + # Replicas are disjoint so there is nothing to coordinate across them, and + # a stateless generation router places locally anyway. + server = _build_adaptor_class()( + config, + req_timeout_secs=req_timeout_secs, + coordinator_url=None, + tokenize_fn=tokenize_fn, + ) + + def _run() -> None: + # OpenAIDisaggServer.__call__ is a coroutine that runs uvicorn, so the + # thread needs its own event loop. Frontend-saturation mitigations + # mirror TRT-LLM's own serve entrypoint: uvloop when available, and + # optionally GC off (tekit measured negligible growth over 200k + # requests; the periodic gen-2 collections otherwise stall the whole + # relay loop). + import os + + try: + import uvloop + + asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) + except ImportError: + pass + if os.environ.get("TRTLLM_DISAGG_SERVER_DISABLE_GC", "0") == "1": + import gc + + gc.disable() + asyncio.run(server(host, port)) + + thread = threading.Thread(target=_run, daemon=True) + thread.start() + + logger.info("TRT-LLM disagg server starting on %s", base_url) + + return thread, base_url, server + + +class DisaggServerActorImpl: + """Hosts one replica's disagg server. + + Its own process rather than an engine worker's: this is the request hot path + for the whole replica, and sharing a process with an engine would couple the + replica's routing latency to that one engine's load -- two uvicorn loops plus + the engine's generate thread pool contending for a single GIL. It holds no + GPUs, so the isolation is cheap. + + Held separately from the ``@ray.remote``-wrapped :class:`DisaggServerActor` + so it can be exercised without Ray. + """ + + def __init__( + self, + replica_idx: int, + frontend_idx: int = 0, + num_frontends: int = 1, + port: int = 0, + serve_args: Optional[dict[str, Any]] = None, + ) -> None: + self._replica_idx = replica_idx + self._frontend_idx = frontend_idx + self._num_frontends = num_frontends + self._port = port + self._thread = None + self._server = None + self._base_url: Optional[str] = None + # Serving from the constructor is what makes crash recovery work: Ray + # re-runs __init__ on actor restart but never replays method calls, so + # with a hard node pin and a deterministic port the restarted frontend + # comes back at the SAME URL and Gym's sticky clients recover after + # their 5xx retries. The empty router trie is a one-time cache-miss. + if serve_args is not None: + self._start_serving(**serve_args) + + def _start_serving( + self, + ctx_addrs: list[tuple[str, int]], + gen_addrs: list[tuple[str, int]], + *, + ctx_router: str, + gen_router: str, + gen_tokids_ctxbytes: bool = False, + gen_strip_message_history: bool = False, + frontend_tokenize: bool = False, + model_name: Optional[str] = None, + default_chat_template_kwargs: Optional[dict[str, Any]] = None, + ) -> None: + # node_id keys the snowflake request-id mint (process_id is hardwired + # to 0 without a coordinator, and time.monotonic() shares an origin + # across processes on one node) -- distinct node_id per frontend is + # what keeps ctx_request_ids disjoint. It also namespaces the router's + # implicit conversation-id counter. + config = build_config( + ctx_addrs, + gen_addrs, + node_id=self._replica_idx * self._num_frontends + self._frontend_idx, + ctx_router=ctx_router, + gen_router=gen_router, + gen_tokids_ctxbytes=gen_tokids_ctxbytes, + gen_strip_message_history=gen_strip_message_history, + ) + tokenize_fn = None + if frontend_tokenize: + assert model_name, "frontend_tokenize requires the model name" + tokenize_fn = _build_frontend_tokenizer( + model_name, default_chat_template_kwargs or {} + ) + self._thread, self._base_url, self._server = start_server( + config, port=self._port, tokenize_fn=tokenize_fn + ) + + def start(self) -> str: + """Wait for the server to answer and return the URL NeMo-Gym will talk to. + + Serving is the constructor's job, not this method's -- that is what + makes Ray actor restart replay it. A second entry point that could also + start would be a silently *less* configured one (no frontend_tokenize, + no model name, hardcoded routers), so the invariant is asserted instead + of being papered over with a fallback. + """ + assert self._base_url is not None, ( + "DisaggServerActorImpl must be constructed with serve_args; the " + "constructor is what starts serving so Ray actor restart replays it." + ) + wait_ready(self._base_url) + # wait_ready cannot distinguish "my server is up" from "a server is up + # on this host:port". uvicorn's bind failure calls sys.exit(1), which + # only unwinds the daemon thread, so without this a port collision is + # answered 200 by the winner and this replica advertises a URL it does + # not own while its engines sit idle. + assert self._thread is not None and self._thread.is_alive(), ( + f"disagg frontend {self._frontend_idx} for replica " + f"{self._replica_idx} died while coming up at {self._base_url} " + "(most likely its port was already bound on this node)" + ) + + logger.info( + "disagg frontend %d/%d for replica %d ready at %s", + self._frontend_idx, + self._num_frontends, + self._replica_idx, + self._base_url, + ) + return self._base_url + + def base_url(self) -> Optional[str]: + return self._base_url + + def shutdown(self) -> bool: + if self._server is not None: + self._server.request_shutdown() + self._server = None + self._thread = None + self._base_url = None + return True + + +def wait_ready(base_url: str, timeout_s: float = 300.0) -> None: + """Block until the disagg server answers /health. + + It reaches out to every engine in its pools on startup, so readiness lags + the thread start by more than a socket bind. + """ + import requests + + health = base_url.rsplit("/v1", 1)[0] + "/health" + deadline = time.monotonic() + timeout_s + while True: + try: + if requests.get(health, timeout=5).status_code == 200: + return + except Exception: + pass + if time.monotonic() > deadline: + raise TimeoutError( + f"disagg server at {base_url} did not become ready within {timeout_s}s" + ) + time.sleep(0.5) + + +@ray.remote(num_cpus=1, num_gpus=0) # pragma: no cover +class DisaggServerActor(DisaggServerActorImpl): + """Ray actor wrapper around :class:`DisaggServerActorImpl`.""" + + pass diff --git a/nemo_rl/models/generation/trtllm/trtllm_generation.py b/nemo_rl/models/generation/trtllm/trtllm_generation.py index 422e2099856..a5759a68cda 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_generation.py +++ b/nemo_rl/models/generation/trtllm/trtllm_generation.py @@ -22,7 +22,7 @@ import asyncio import os from collections import defaultdict -from typing import Any, AsyncGenerator, Optional, Union, cast +from typing import Any, AsyncGenerator, Literal, Optional, Union, cast import numpy as np import ray @@ -37,7 +37,11 @@ GenerationOutputSpec, reject_unenforceable_refit_deadline, ) -from nemo_rl.models.generation.trtllm.config import TrtllmConfig +from nemo_rl.models.generation.trtllm.config import ( + TrtllmConfig, + TrtllmDisaggConfig, + resolve_trtllm_disagg_config, +) class TrtllmGeneration(GenerationInterface): @@ -50,16 +54,49 @@ def init_cluster_placement_groups( ) -> None: """Pre-initialize placement groups matching TRT-LLM's topology.""" trtllm_cfg = config["trtllm_cfg"] - tp = trtllm_cfg["tensor_parallel_size"] + disagg = resolve_trtllm_disagg_config(config) + engine_tp = trtllm_cfg["tensor_parallel_size"] + if disagg.enabled: + engine_tp = max( + int( + disagg.role_trtllm_kwargs(role).get( + "tensor_parallel_size", engine_tp + ) + ) + for role in ("ctx", "gen") + ) pp = trtllm_cfg.get("pipeline_parallel_size", 1) assert pp == 1, ( "TRT-LLM backend does not support pipeline parallelism yet " f"(pipeline_parallel_size={pp}, must be 1)." ) - model_parallel_size = tp * pp + # GPUs held by the *widest single engine*. Deliberately not a replica's + # width: this decides use_unified_pg, which exists so one tied worker + # group can take bundles across nodes, and the tied group is an engine + # (_get_tied_worker_bundle_indices slices per engine). A replica is not + # a scheduling unit -- its engines are independent actor groups talking + # over HTTP and the KV transceiver, so it may span nodes, and we only + # soft-pin it for locality. Sizing this by the replica would push + # node-local layouts (e.g. 4 x TP2 engines on 4-GPU nodes) into a + # unified PG for nothing. + widest_engine_gpus = engine_tp * pp colocated = bool(config.get("colocated", {}).get("enabled", False)) - needs_cross_node = model_parallel_size > cluster.num_gpus_per_node + # Colocated time-multiplexes the GPUs: engines sleep (dropping the whole + # KV pool) while the policy trains. Disaggregation cannot survive that -- + # a replica's context and generation engines must be resident *at the + # same time* for the transceiver to hand a prefilled cache over, and the + # disagg servers hold HTTP connections to engines that would be asleep. + # Reject the combination instead of hanging on the first request after a + # sleep. + assert not (disagg.enabled and colocated), ( + "PD disaggregation requires non-colocated generation: colocated mode " + "sleeps the engines between rollouts, which drops the KV cache the " + "transceiver needs. Set colocated.enabled=false or " + "trtllm_cfg.disaggregation.enabled=false." + ) + + needs_cross_node = widest_engine_gpus > cluster.num_gpus_per_node assert not (needs_cross_node and colocated), ( "TRT-LLM cross-node tensor parallelism is only supported for " "non-colocated generation." @@ -79,26 +116,61 @@ def __init__( ): self.cfg = config self.tp_size = self.cfg["trtllm_cfg"]["tensor_parallel_size"] - self.model_parallel_size = self.tp_size - - assert cluster.world_size() % self.model_parallel_size == 0, ( - f"Cluster world_size ({cluster.world_size()}) must be divisible by " - f"TP size ({self.model_parallel_size})." - ) - self.dp_size = cluster.world_size() // self.model_parallel_size - - # MoE: TRT-LLM partitions TP on MoE layers into moe_tp × moe_ep, so - # the product must equal the main tensor_parallel_size. Validate here - # to fail fast — the LLM constructor would otherwise raise a less - # actionable error deep inside the engine. - moe_tp = self.cfg["trtllm_cfg"].get("moe_tensor_parallel_size") - moe_ep = self.cfg["trtllm_cfg"].get("moe_expert_parallel_size") - if moe_tp is not None or moe_ep is not None: - moe_tp_v = moe_tp if moe_tp is not None else 1 - moe_ep_v = moe_ep if moe_ep is not None else 1 - assert moe_tp_v * moe_ep_v == self.tp_size, ( - f"moe_tensor_parallel_size ({moe_tp_v}) * moe_expert_parallel_size " - f"({moe_ep_v}) must equal tensor_parallel_size ({self.tp_size})." + # Validated once here rather than per access: every disagg default + # lives on the schema, so the rest of this class reads attributes + # instead of re-deriving a default per key. + self._disagg = resolve_trtllm_disagg_config(config) + + # Per-engine role and TP width, in engine order -- the single source of + # truth for how the cluster is sliced. Without disaggregation every + # engine is identical; under it, each replica contributes its context + # engines followed by its generation engines. + ( + self._engine_roles, + self._engine_tps, + self.num_replicas, + ) = self._plan_engines(cluster.world_size()) + # DP keeps its general meaning -- the number of independent rollout + # shards, i.e. replicas. It is deliberately *not* the engine count: + # under disaggregation a replica is several engines, and overloading + # "dp" to mean engines is what makes the two concepts collide. + self.dp_size = self.num_replicas + self.num_engines = len(self._engine_tps) + + # GPUs held by the widest single engine; see + # init_cluster_placement_groups for why this is per engine, not per + # replica. + self.widest_engine_gpus = max(self._engine_tps) + + assert sum(self._engine_tps) == cluster.world_size(), ( + f"Engine layout needs {sum(self._engine_tps)} GPUs " + f"({self._engine_tps}) but the cluster has {cluster.world_size()}." + ) + + # MoE: TRT-LLM partitions TP on MoE layers into moe_tp x moe_ep, so the + # product must equal that engine's TP width. Validate here to fail fast + # -- the LLM constructor would otherwise raise a less actionable error + # deep inside the engine. Under disaggregation this is per role, since + # both TP and the MoE split can differ between prefill and decode. + if self._disagg_cfg.enabled: + engine_configs = [ + (f"{role}_trtllm_kwargs", self._role_kwargs(role)) + for role in ("ctx", "gen") + ] + else: + engine_configs = [("trtllm_cfg", self.cfg["trtllm_cfg"])] + + for label, kwargs in engine_configs: + m_tp = kwargs.get("moe_tensor_parallel_size") + m_ep = kwargs.get("moe_expert_parallel_size") + if m_tp is None and m_ep is None: + continue + product = (m_tp or 1) * (m_ep or 1) + engine_tp = int(kwargs["tensor_parallel_size"]) + assert product == engine_tp, ( + f"{label}: moe_tensor_parallel_size ({m_tp}) * " + f"moe_expert_parallel_size ({m_ep}) = {product} must equal " + f"tensor_parallel_size ({engine_tp})." ) missing_keys = [k for k in TrtllmConfig.__required_keys__ if k not in self.cfg] @@ -106,15 +178,31 @@ def __init__( missing_keys.append("model_name") assert not missing_keys, f"TrtllmConfig missing keys: {missing_keys}" + # [data_parallel, tensor_parallel] = [replicas, GPUs per replica]. A + # replica is the unit data is actually parallelised over -- one rollout + # shard, one URL -- and its GPU span tiles the cluster exactly, so the + # grid is rectangular by construction with no special case for + # asymmetric per-role TP. Without disaggregation a replica is a single + # engine, so this is the usual arange(world).reshape(dp, tp). + # + # Under disaggregation the second axis spans a whole replica, not one + # engine's TP group, so get_axis_size("tensor_parallel") is not any + # engine's TP width -- read self._engine_tps for that. Nothing + # dispatches on the axis: per-engine fan-out goes through + # _run_on_engines, and the two grid-driven worker-group calls sit + # behind _assert_direct_dispatch_allowed, which disaggregation rejects. + world_size = sum(self._engine_tps) self.sharding_annotations = NamedSharding( - layout=np.arange(cluster.world_size()).reshape( - self.dp_size, - self.tp_size, - ), + layout=np.arange(world_size).reshape(self.num_replicas, -1), names=["data_parallel", "tensor_parallel"], ) self.colocated_enabled = bool(self.cfg["colocated"]["enabled"]) + self.async_engine = bool(self.cfg["trtllm_cfg"].get("async_engine", False)) + # One disagg server actor and one URL per replica; stay empty unless + # disaggregation is on. See _start_disagg_servers. + self._disagg_actors: list[ray.actor.ActorHandle] = [] + self._disagg_server_urls: list[Optional[str]] = [] # The synchronous TRT-LLM engine path is no longer supported: only the # async worker wires up colocated sleep/wakeup, IPC-ZMQ refit, and # per-sample streaming. Fail loudly at setup rather than silently @@ -126,8 +214,53 @@ def __init__( self.init_cluster_placement_groups(cluster, config) + # Engines differ from one another only under disaggregation, so the + # explicit bundle list (which fixes engine order) is only required + # there; a uniform run keeps the original workers_per_node path. + use_explicit_bundles = self.widest_engine_gpus > 1 or self._disagg_cfg.enabled + node_bundle_indices = ( + self._get_tied_worker_bundle_indices(cluster) + if use_explicit_bundles + else None + ) + # Global rank of each engine's model owner, in engine order -- the list + # every per-engine call is dispatched on (see _run_on_engines). The + # worker group hands bundle_indices (the "you build the AsyncLLM" + # signal) to each tied group's local rank 0 and assigns ranks in bundle + # order, so the owners are the prefix sums of the tied group widths. + # Derived from the very bundle split handed to the worker group so + # engine fan-out cannot drift from it; the worker group's DP leaders are + # a strict subset of these (one per replica) and stop coinciding with + # them under disaggregation, which is why engine fan-out keeps its own + # list. + engine_widths = ( + [len(indices) for _, indices in node_bundle_indices] + if node_bundle_indices is not None + # No explicit bundles means every engine is one worker of its own. + else [1] * self.num_engines + ) + assert engine_widths == self._engine_tps, ( + f"Bundle split gives engine widths {engine_widths} but the engine " + f"layout expects {self._engine_tps}; every per-engine kwarg would " + "land on the wrong engine." + ) + self._engine_owner_indices: list[int] = [] + offset = 0 + for width in engine_widths: + self._engine_owner_indices.append(offset) + offset += width + + # Owner of each replica's first engine: the DP leaders the worker group + # would otherwise derive per engine, which is only right while a replica + # is a single engine. + self._replica_owner_indices = self._engine_owner_indices[ + :: self.num_engines // self.num_replicas + ] + worker_cls = "nemo_rl.models.generation.trtllm.trtllm_worker_async.TrtllmAsyncGenerationWorker" - worker_builder = RayWorkerBuilder(worker_cls, config) + worker_builder = RayWorkerBuilder( + worker_cls, self._config_with_engine_overrides(node_bundle_indices) + ) # NCCL_CUMEM_ENABLE=1 is needed for the non-colocated NCCL collective # broadcast; colocated shares the policy's NCCL group so don't touch it. @@ -135,8 +268,7 @@ def __init__( if not self.colocated_enabled: env_vars["NCCL_CUMEM_ENABLE"] = "1" - if self.model_parallel_size > 1: - node_bundle_indices = self._get_tied_worker_bundle_indices(cluster) + if node_bundle_indices is not None: self.worker_group = RayWorkerGroup( cluster, worker_builder, @@ -144,6 +276,7 @@ def __init__( bundle_indices_list=node_bundle_indices, sharding_annotations=self.sharding_annotations, env_vars=env_vars, + dp_leader_worker_indices=self._replica_owner_indices, ) else: self.worker_group = RayWorkerGroup( @@ -158,9 +291,8 @@ def __init__( # post-init on workers (starts HTTP server when expose_http_server=true, # finishes async engine setup for the async worker variant). post_init_method = "post_init_async" - futures = self.worker_group.run_all_workers_single_data( + futures = self._run_on_engines( post_init_method, - run_rank_0_only_axes=["tensor_parallel"], ) ray.get(futures) @@ -171,10 +303,132 @@ def __init__( self.device_uuids = self._report_device_id() + # DP is replicas, not engines: only equal when a replica is one engine. assert self.dp_size == self.worker_group.dp_size, ( - f"DP size mismatch: expected {self.dp_size}, got {self.worker_group.dp_size}" + f"Replica count {self.dp_size} does not match the worker group's DP " + f"size {self.worker_group.dp_size}." + ) + + # ------------------------------------------------------------------ # + # Engine layout + # ------------------------------------------------------------------ # + + def _role_kwargs(self, role: Literal["ctx", "gen"]) -> dict[str, Any]: + """This role's engine overrides, merged over the base ``trtllm_cfg``. + + Any TRT-LLM kwarg may be overridden per role; TP and the MoE split are + the ones that usually differ between prefill and decode. + """ + overrides = self._disagg_cfg.role_trtllm_kwargs(role) + return {**self.cfg["trtllm_cfg"], **overrides} + + def _plan_engines(self, world_size: int) -> tuple[list[str], list[int], int]: + r"""Per-engine ``(role, tp_width)``, in engine order. + + Without disaggregation every engine is a plain generation engine of + ``tensor_parallel_size`` GPUs. Under disaggregation each replica + contributes its context engines followed by its generation engines: + + [ctx_tp x M, gen_tp x K, ctx_tp x M, gen_tp x K, ...] + \______ replica 0 _____/ \____ replica 1 ... + + The replica *count* is derived, not configured: it follows from the + cluster size, the same way the DP-shard count does without + disaggregation. Same-role engines stay contiguous so each role's GPUs + are adjacent. + """ + disagg = self._disagg_cfg + if not disagg.enabled: + assert world_size % self.tp_size == 0, ( + f"Cluster world_size ({world_size}) must be divisible by " + f"TP size ({self.tp_size})." + ) + n = world_size // self.tp_size + # Without disaggregation a replica *is* an engine: one DP shard, + # one URL, nothing below it to front. Reporting n rather than 1 + # keeps "replica" meaning the same thing on both paths. + return ["generation"] * n, [self.tp_size] * n, n + + ctx_tp = int(self._role_kwargs("ctx")["tensor_parallel_size"]) + gen_tp = int(self._role_kwargs("gen")["tensor_parallel_size"]) + for role, value in (("ctx", ctx_tp), ("gen", gen_tp)): + assert value >= 1, f"{role}_trtllm_kwargs.tensor_parallel_size must be >= 1" + + num_ctx = disagg.num_context_engines + num_gen = disagg.num_generation_engines + assert num_ctx >= 1 and num_gen >= 1, ( + f"a replica needs at least one engine of each role, got " + f"num_context_engines={num_ctx}, num_generation_engines={num_gen}" ) + replica_width = num_ctx * ctx_tp + num_gen * gen_tp + assert world_size % replica_width == 0, ( + f"replica width {replica_width} GPUs " + f"({num_ctx} ctx x TP{ctx_tp} + {num_gen} gen x TP{gen_tp}) does not " + f"divide the {world_size} inference GPUs." + ) + num_replicas = world_size // replica_width + + roles = (["context"] * num_ctx + ["generation"] * num_gen) * num_replicas + tps = ([ctx_tp] * num_ctx + [gen_tp] * num_gen) * num_replicas + return roles, tps, num_replicas + + def _config_with_engine_overrides( + self, node_bundle_indices: Optional[list[tuple[int, list[int]]]] + ) -> TrtllmConfig: + """Attach each engine's role and construction overrides to the worker config. + + Engine parallelism is fixed when ``AsyncLLM`` is built, and a worker + cannot infer its own values from its TP width -- two roles may share a + width and still want different expert layouts. ``RayWorkerBuilder`` hands + every worker the same config, so the driver attaches a map keyed by + :meth:`_engine_key` and each worker looks up its own entry. + + The context/generation *role* is per request in TRT-LLM, so an engine + does not strictly need to know it; it is recorded anyway because the + transceiver config and the role's kwargs are chosen from it. + """ + if node_bundle_indices is None or not self._disagg_cfg.enabled: + return self.cfg + + overrides: dict[str, dict[str, Any]] = {} + role_counts: dict[str, int] = {} + for (pg_idx, bundles), role in zip( + node_bundle_indices, self._engine_roles, strict=True + ): + prefix: Literal["ctx", "gen"] = "ctx" if role == "context" else "gen" + # Ordinal within the role, so a layout with several engines of one + # role (CTX_ENGINES=2, or more than one replica) can still tell them + # apart. Only consumers that need a stable per-engine name use it -- + # the nsys report filename, so far. + ordinal = role_counts.get(role, 0) + role_counts[role] = ordinal + 1 + # Every engine gets an entry: the worker treats a missing key as a + # driver/worker mismatch, so absence must never read as "no + # overrides for this engine". + overrides[self._engine_key(pg_idx, bundles)] = { + "_disagg_role": role, + "_disagg_role_ordinal": ordinal, + **self._disagg_cfg.role_trtllm_kwargs(prefix), + } + + cfg = dict(self.cfg) + cfg["trtllm_cfg"] = {**self.cfg["trtllm_cfg"], "_engine_overrides": overrides} + return cast(TrtllmConfig, cfg) + + @staticmethod + def _engine_key(pg_idx: int, local_bundle_indices: list[int]) -> str: + """Stable id for the engine occupying these bundles. + + Both sides of the worker boundary compute this from the same + ``(pg_idx, local_bundle_indices)`` tuple, so the worker can look up its + own per-engine overrides without the driver needing per-worker init + kwargs. Keying on the tuple rather than deriving a global ordinal keeps + it correct for both the unified-PG and per-node-PG layouts, whose local + bundle indices mean different things. + """ + return f"{pg_idx}:" + ",".join(str(i) for i in local_bundle_indices) + # ------------------------------------------------------------------ # # Placement helpers (simplified from VllmGeneration) # ------------------------------------------------------------------ # @@ -189,12 +443,16 @@ def _get_tied_worker_bundle_indices( per-node placement groups (node-local model parallelism). For unified PGs, bundles are reordered by physical node before slicing so each TP group stays as node-local as possible. + + Bundles are consumed engine by engine following ``self._engine_tps``, so + engines of differing width (PD disaggregation with asymmetric TP) each + get exactly their own number of bundles. """ placement_groups = cluster.get_placement_groups() if not placement_groups: raise ValueError("No placement groups available in the cluster") - model_parallel_size = self.model_parallel_size + engine_tps = self._engine_tps if len(placement_groups) == 1: # Single unified PG: TP > GPUs/node, so model parallelism may span @@ -222,13 +480,6 @@ def _get_tied_worker_bundle_indices( counts = [len(b) for b in node_bundles.values()] assert len(set(counts)) == 1, "All nodes must have identical bundle counts" - total = sum(counts) - num_groups = total // model_parallel_size - if num_groups == 0: - raise ValueError( - "Unable to allocate any worker groups with the available resources." - ) - # RayVirtualCluster records the physical-node bundle order when it # builds a unified PG. Preserve it so TP replicas occupy contiguous # nodes in the topology-aware order selected by the cluster. @@ -237,23 +488,52 @@ def _get_tied_worker_bundle_indices( for nid in sorted(node_bundles): flat.extend(node_bundles[nid]) + if len(flat) < sum(engine_tps): + raise ValueError( + f"Engine layout needs {sum(engine_tps)} bundles but the " + f"unified placement group has {len(flat)}." + ) + tied_groups: list[tuple[int, list[int]]] = [] - for i in range(num_groups): - slice_ = flat[i * model_parallel_size : (i + 1) * model_parallel_size] + cursor = 0 + for tp in engine_tps: # The first value is a placement-group index. # A unified cluster has exactly one PG (index 0). - tied_groups.append((0, slice_)) + tied_groups.append((0, flat[cursor : cursor + tp])) + cursor += tp else: tied_groups = [] - for pg_idx, pg in enumerate(placement_groups): + engine_idx = 0 + # Consume placement groups in topology order, not creation order. Ray + # picks the physical node behind each per-node PG, so consecutive + # pg_idx values can land in different NVLink domains -- and engines + # are laid out replica by replica, which would put a replica's + # context and generation engines on opposite sides of the fabric and + # push its KV handoff onto InfiniBand. Only the *order* changes; + # every PG is still used exactly once. + for pg_idx in cluster.get_topology_sorted_pg_indices(): + pg = placement_groups[pg_idx] if pg.bundle_count == 0: continue - num_groups_in_pg = pg.bundle_count // model_parallel_size - for group_idx in range(num_groups_in_pg): - start_idx = group_idx * model_parallel_size - end_idx = start_idx + model_parallel_size - bundle_indices = list(range(start_idx, end_idx)) - tied_groups.append((pg_idx, bundle_indices)) + cursor = 0 + while engine_idx < len(engine_tps): + tp = engine_tps[engine_idx] + if cursor + tp > pg.bundle_count: + break + tied_groups.append((pg_idx, list(range(cursor, cursor + tp)))) + cursor += tp + engine_idx += 1 + if cursor != pg.bundle_count: + # An engine may not straddle a placement group (== node): + # TRT-LLM does not support cross-node TP, and a partially + # filled node would leave GPUs idle while the engine count + # silently drops. + raise ValueError( + f"Engine widths {engine_tps} do not tile placement group " + f"{pg_idx} ({pg.bundle_count} bundles): {cursor} bundles " + f"used. Choose per-role TP sizes whose group total " + f"divides the GPUs per node." + ) if not tied_groups: raise ValueError( @@ -261,21 +541,214 @@ def _get_tied_worker_bundle_indices( ) return tied_groups + # ------------------------------------------------------------------ # + # PD disaggregation + # ------------------------------------------------------------------ # + + @property + def _disagg_cfg(self) -> TrtllmDisaggConfig: + return self._disagg + + def _assert_direct_dispatch_allowed(self) -> None: + """Reject the token-in-token-out path while PD is enabled. + + ``generate`` / ``generate_async`` round-robin straight to DP leaders, + bypassing the group entry points. That would still produce correct + tokens -- each engine can prefill and decode on its own -- but it would + run *without* disaggregation while the user believes they are + exercising it. Fail loudly instead. + """ + if self._disagg_cfg.enabled: + raise RuntimeError( + "PD disaggregation is only wired for the HTTP/NeMo-Gym rollout " + "path; TrtllmGeneration.generate()/generate_async() dispatch " + "directly to engines and would silently bypass it. Use the " + "NeMo-Gym entrypoint, or set " + "trtllm_cfg.disaggregation.enabled=false." + ) + + def _start_disagg_servers(self) -> list[Optional[str]]: + """Start one disagg server per replica and return their URLs. + + Every engine already runs its own HTTP server; an address is the only + thing a disagg server is given about one. This slices the engine list + into replicas, hands each disagg server its two address pools, and + collects the single URL it exposes. Those URLs are what NeMo-Gym sees. + + Starting a server is not idempotent, so a second call returns the URLs + of the servers already running instead of doubling them. + """ + if self._disagg_server_urls: + return self._disagg_server_urls + + disagg = self._disagg_cfg + num_ctx = disagg.num_context_engines + num_gen = disagg.num_generation_engines + per_replica = num_ctx + num_gen + + assert self.cfg["trtllm_cfg"].get("expose_http_server"), ( + "PD disaggregation requires trtllm_cfg.expose_http_server=true: an " + "address is the only way the disagg server can reach an engine." + ) + + reported = self._report_engine_addrs() + missing = [i for i, a in enumerate(reported) if not a] + assert not missing, f"engines {missing} reported no HTTP address" + # Rebind after the guard so the element type is non-optional from here + # down; asserting on a separate list does not narrow `reported` itself. + addrs: list[dict[str, Any]] = [a for a in reported if a] + + from ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy + + from nemo_rl.models.generation.trtllm.trtllm_disagg_server import ( + DisaggServerActor, + ) + + n_fe = disagg.num_frontend_workers + assert self.num_replicas * n_fe <= 256, ( + "replicas * num_frontend_workers must fit the snowflake node_id " + f"space (8 bits): {self.num_replicas} * {n_fe} > 256" + ) + # Deterministic ports: a restarted frontend re-binds the same port on + # its pinned node, so the URL Gym holds stays valid across crashes. + # Keep the base outside virtual_cluster's random master-port window + # (1400-1999). The offset must carry BOTH indices: several replicas can + # land on one node (a node holds `gpus_per_node / replica_width` + # replicas), and with only frontend_idx in it every replica's frontend + # 0 would ask for the same port -- the losers die inside their daemon + # thread while /health is answered 200 by the winner bound on 0.0.0.0, + # so they silently advertise the winner's URL and their engines idle. + base_port = disagg.frontend_base_port + + self._disagg_actors = [] + futures = [] + for replica_idx in range(self.num_replicas): + base = replica_idx * per_replica + serve_args = dict( + ctx_addrs=[ + (a["host"], a["port"]) for a in addrs[base : base + num_ctx] + ], + gen_addrs=[ + (a["host"], a["port"]) + for a in addrs[base + num_ctx : base + per_replica] + ], + ctx_router=disagg.ctx_router, + gen_router=disagg.gen_router, + gen_tokids_ctxbytes=disagg.gen_tokids_ctxbytes, + gen_strip_message_history=disagg.gen_strip_message_history, + frontend_tokenize=disagg.frontend_tokenize, + model_name=self.cfg["model_name"], + default_chat_template_kwargs=self.cfg["trtllm_cfg"].get( + "default_chat_template_kwargs" + ), + ) + for fe_idx in range(n_fe): + # Its own CPU-only actor rather than an engine worker's + # process: this is the request hot path for the whole replica, + # and sharing a process with an engine would couple the + # replica's routing latency to that one engine's load. + # Frontends spread round-robin across the replica's engine + # nodes; the pin is HARD and restarts are unlimited so a + # crashed frontend re-serves the same host:port (the + # constructor starts the server -- Ray never replays method + # calls on restart). + pin_node = addrs[base + (fe_idx % per_replica)]["node_id"] + actor = DisaggServerActor.options( + scheduling_strategy=NodeAffinitySchedulingStrategy( + node_id=pin_node, soft=False + ), + name=f"trtllm_disagg_server_{replica_idx}_{fe_idx}", + max_restarts=-1, + # The server imports tensorrt_llm (OpenAIDisaggServer, + # disagg_utils), which only exists in the engine workers' + # venv. Without this the actor starts in the driver's + # environment and dies with ModuleNotFoundError. + runtime_env={"py_executable": self.worker_group.py_executable}, + ).remote( + replica_idx, + frontend_idx=fe_idx, + num_frontends=n_fe, + port=base_port + replica_idx * n_fe + fe_idx, + serve_args=serve_args, + ) + self._disagg_actors.append(actor) + # start() waits for readiness and returns the URL; the server + # itself was already brought up by the constructor. + futures.append(actor.start.remote()) + + self._disagg_server_urls = ray.get(futures) + print( + f" ✓ PD disaggregation: {self.num_replicas} replica(s) x " + f"({num_ctx} context + {num_gen} generation) engines, " + f"{n_fe} frontend worker(s) each; " + f"disagg servers: {self._disagg_server_urls}", + flush=True, + ) + return self._disagg_server_urls + + def _run_on_engines( + self, + method_name: str, + per_engine: Optional[dict[str, list[Any]]] = None, + **common: Any, + ) -> list[ray.ObjectRef]: + """Fan out to each engine's model owner, one call per engine. + + Dispatches by :attr:`_engine_owner_indices`, computed once from the + bundle split. Neither the sharding grid's ``run_rank_0_only_axes`` gate + nor the worker group's DP-leader list is used: the grid cannot model + engines of differing width, and the DP leaders are one worker per + *replica*, which under disaggregation skips every engine but the + replica's first. This keeps one source of truth for who owns an engine. + + Args: + per_engine: kwargs whose value is a per-engine list, in engine order. + **common: kwargs passed unchanged to every engine. + """ + futures = [] + for engine_idx, worker_idx in enumerate(self._engine_owner_indices): + kwargs = dict(common) + for key, values in (per_engine or {}).items(): + kwargs[key] = values[engine_idx] + futures.append( + self.worker_group.run_single_worker_single_data( + method_name=method_name, worker_idx=worker_idx, **kwargs + ) + ) + return futures + + def _report_engine_addrs(self) -> list[Optional[dict[str, Any]]]: + """Collect ``{host, port, node_id}`` from every engine's HTTP server. + + ``None`` for an engine whose server never started. The node id is what + lets the replica's disagg server be pinned beside its engines. + """ + futures = self._run_on_engines( + "report_http_addr", + ) + return ray.get(futures) + def _report_dp_openai_server_base_urls(self) -> list[Optional[str]]: - """Collect HTTP server base URLs from each DP-rank-0 worker.""" + """One or more OpenAI-compatible base URLs per DP shard, in shard order. + + A shard's entry point is not always an engine: under disaggregation the + caller must talk to the replica's disagg server, which routes prefill and + decode itself. Branching here rather than at the call site keeps the + contract "one URL per DP shard" true on both paths -- an engine-per-URL + list would have ``num_engines`` entries, which is no longer ``dp_size``. + """ + if self._disagg_cfg.enabled: + return self._start_disagg_servers() if not self.cfg["trtllm_cfg"].get("expose_http_server"): - urls = [cast(Optional[str], None)] * self.dp_size - return urls - futures = self.worker_group.run_all_workers_single_data( + return [cast(Optional[str], None)] * self.dp_size + futures = self._run_on_engines( "report_dp_openai_server_base_url", - run_rank_0_only_axes=["tensor_parallel"], ) return ray.get(futures) def _report_device_id(self) -> list[list[str]]: - futures = self.worker_group.run_all_workers_single_data( + futures = self._run_on_engines( "report_device_id_async", - run_rank_0_only_axes=["tensor_parallel"], ) return ray.get(futures) @@ -294,20 +767,28 @@ def init_collective( if not self.worker_group or not self.worker_group.workers: raise RuntimeError("Worker group not initialised") - total_workers = len(self.worker_group.workers) - workers_per_group = total_workers // self.dp_size - rank_prefix_list = list(range(0, total_workers, workers_per_group)) + # The caller sizes the group from the cluster config while the ranks + # below come from the engine layout. If the two disagree the group never + # fills up and StatelessProcessGroup's TCPStore blocks forever, so check + # it here rather than debugging a hang at startup. + inference_world_size = sum(self._engine_tps) + assert train_world_size + inference_world_size == world_size, ( + f"Collective world_size {world_size} != train {train_world_size} + " + f"inference {inference_world_size} (engine widths {self._engine_tps})." + ) - return self.worker_group.run_all_workers_multiple_data( + # rank_prefix is each engine's first rank inside the inference half of + # the group; the engine's own TP ranks follow contiguously. Not a fixed + # stride -- with asymmetric TP the engines are not equally wide, and a + # wrong offset would map generation ranks onto the wrong slice of the + # training-side broadcast, silently loading the wrong weights. + return self._run_on_engines( "init_collective_async", - rank_prefix=rank_prefix_list, - run_rank_0_only_axes=["tensor_parallel"], - common_kwargs={ - "ip": ip, - "port": port, - "world_size": world_size, - "train_world_size": train_world_size, - }, + per_engine={"rank_prefix": self._engine_owner_indices}, + ip=ip, + port=port, + world_size=world_size, + train_world_size=train_world_size, ) def generate( @@ -315,10 +796,11 @@ def generate( data: BatchedDataDict[GenerationDatumSpec], greedy: bool = False, ) -> BatchedDataDict[GenerationOutputSpec]: + self._assert_direct_dispatch_allowed() assert isinstance(data, BatchedDataDict) assert "input_ids" in data and "input_lengths" in data - dp_size = self.sharding_annotations.get_axis_size("data_parallel") + dp_size = self.dp_size sharded_data = cast( list[SlicedDataDict], data.shard_by_batch_size( @@ -364,6 +846,7 @@ async def generate_async( in-flight Ray calls share the same AsyncLLM, which batches them internally via asyncio.gather. """ + self._assert_direct_dispatch_allowed() if "input_ids" not in data or "input_lengths" not in data: raise AssertionError( "input_ids and input_lengths are required in data for generate_async" @@ -374,9 +857,9 @@ async def generate_async( f"generate_async expects single-sample data, got batch_size={data.size}." ) - leader_worker_idx = self.worker_group.get_dp_leader_worker_idx( + leader_worker_idx = self._engine_owner_indices[ self.current_generate_dp_shard_idx - ) + ] worker_result_ref = self.worker_group.run_single_worker_single_data( method_name="generate_async", worker_idx=leader_worker_idx, @@ -385,7 +868,7 @@ async def generate_async( ) self.current_generate_dp_shard_idx = ( self.current_generate_dp_shard_idx + 1 - ) % self.worker_group.dp_size + ) % self.num_engines timeout_seconds = float( os.environ.get("NRL_TRTLLM_ASYNC_TIMEOUT_SECONDS", "900") @@ -408,9 +891,8 @@ def prepare_for_generation(self, *args: Any, **kwargs: Any) -> bool: if not self.colocated_enabled: return True try: - futures = self.worker_group.run_all_workers_single_data( + futures = self._run_on_engines( "wake_up_async", - run_rank_0_only_axes=["tensor_parallel"], **kwargs, ) results = ray.get(futures) @@ -426,9 +908,8 @@ def finish_generation(self, *args: Any, **kwargs: Any) -> bool: method_name = "sleep_async" else: method_name = "reset_prefix_cache_async" - futures = self.worker_group.run_all_workers_single_data( + futures = self._run_on_engines( method_name, - run_rank_0_only_axes=["tensor_parallel"], ) results = ray.get(futures) return all(r for r in results if r is not None) @@ -437,10 +918,9 @@ def finish_generation(self, *args: Any, **kwargs: Any) -> bool: return False def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: - futures = self.worker_group.run_all_workers_single_data( + futures = self._run_on_engines( "prepare_refit_info_async", state_dict_info=state_dict_info, - run_rank_0_only_axes=["tensor_parallel"], ) ray.get(futures) @@ -448,9 +928,8 @@ def start_gpu_profiling(self) -> None: """Grpo profiling protocol: start nsys capture on the GPU workers.""" if not self.worker_group or not self.worker_group.workers: return - futures = self.worker_group.run_all_workers_single_data( + futures = self._run_on_engines( "start_gpu_profiling_async" if self.async_engine else "start_gpu_profiling", - run_rank_0_only_axes=["tensor_parallel"], ) ray.get(futures) @@ -458,9 +937,8 @@ def stop_gpu_profiling(self) -> None: """Grpo profiling protocol: stop nsys capture on the GPU workers.""" if not self.worker_group or not self.worker_group.workers: return - futures = self.worker_group.run_all_workers_single_data( + futures = self._run_on_engines( "stop_gpu_profiling_async" if self.async_engine else "stop_gpu_profiling", - run_rank_0_only_axes=["tensor_parallel"], ) ray.get(futures) @@ -473,9 +951,8 @@ def update_weights_from_collective( trtllm_cfg = self.cfg["trtllm_cfg"] in_flight = bool(trtllm_cfg.get("in_flight_weight_updates")) recompute_kv = bool(trtllm_cfg.get("recompute_kv_cache_after_weight_updates")) - return self.worker_group.run_all_workers_single_data( + return self._run_on_engines( "update_weights_from_collective_async", - run_rank_0_only_axes=["tensor_parallel"], drain=not in_flight, recompute_kv=recompute_kv, ) @@ -484,9 +961,8 @@ def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: """Receive weights via CUDA-IPC + ZMQ (colocated mode).""" if not self.worker_group or not self.worker_group.workers: raise RuntimeError("Worker group not initialised") - return self.worker_group.run_all_workers_single_data( + return self._run_on_engines( "update_weights_via_ipc_zmq_async", - run_rank_0_only_axes=["tensor_parallel"], ) def invalidate_kv_cache(self) -> bool: @@ -532,6 +1008,17 @@ def get_logger_metrics(self) -> dict[str, Any]: def shutdown(self) -> bool: try: + # Stop routing before the engines go away, so in-flight requests are + # not handed to an endpoint that is already tearing down. + for actor in getattr(self, "_disagg_actors", []): + try: + ray.get(actor.shutdown.remote(), timeout=30) + except Exception as e: + print(f"Error stopping disagg server: {e}") + ray.kill(actor) + self._disagg_actors = [] + self._disagg_server_urls = [] + return self.worker_group.shutdown(cleanup_method="shutdown") except Exception as e: print(f"Error during TRT-LLM shutdown: {e}") diff --git a/nemo_rl/models/generation/trtllm/trtllm_http_server.py b/nemo_rl/models/generation/trtllm/trtllm_http_server.py index bf40dd78143..fde2c2840bb 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_http_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_http_server.py @@ -13,11 +13,24 @@ # limitations under the License. """OpenAI-compatible HTTP server wrapping ``tensorrt_llm.LLM``, serving /v1/chat/completions. -Returns NeMoGym fields (prompt_token_ids, generation_token_ids, generation_log_probs). -Supports Qwen3 tool calling, DeepSeekR1Parser reasoning, and prefix token splicing. +Returns prompt and generated token ids alongside per-token logprobs, and supports +Qwen3 tool calling, DeepSeekR1Parser reasoning, and prefix token splicing. + +Under PD disaggregation this endpoint is *leg-aware*. A replica's +``OpenAIDisaggServer`` drives it twice per request: + +* ``context_only`` -- prefill only. Returns the handshake and the prompt token + ids, skipping all post-processing; see :func:`_context_leg_response`. +* ``generation_only`` -- decodes and post-processes as usual, but takes the + prompt token ids the orchestrator relays from the context leg rather than + rebuilding them, so the sequence matches the KV that was transferred. """ +import asyncio +import functools import logging +import os +import random import threading import time import uuid @@ -33,6 +46,152 @@ logger = logging.getLogger(__name__) +@functools.cache +def _choice_declares_token_ids() -> bool: + """Whether the installed TRT-LLM's chat choice model declares ``token_ids``. + + Cached because it is consulted once per generated response, and imported + lazily because this module is also used on the aggregated path where + tensorrt_llm need not be importable. + """ + try: + from tensorrt_llm.serve.openai_protocol import ChatCompletionResponseChoice + except ImportError: + return False + return "token_ids" in ChatCompletionResponseChoice.model_fields + + +def _tokenizer_backend_name(tokenizer: Any) -> str: + """Return the concrete backend that performs encode/decode operations.""" + backend = getattr(tokenizer, "_tokenizer", None) + implementation = backend if backend is not None else tokenizer + implementation_type = type(implementation) + return f"{implementation_type.__module__}.{implementation_type.__name__}" + + +async def build_spliced_prompt_ids( + messages: list[dict], + tools: "list[dict] | None", + tokenizer: Any, + model_config: Any, + template_kwargs: dict[str, Any], +) -> list[int]: + """Render the chat template and splice in the on-policy prefix. + + The single source of truth for turning a conversation into engine prompt + ids -- shared by the engine adapter's route and the disagg frontend's + tokenizer (phase-2 ``frontend_tokenize``) so both produce IDENTICAL ids; + a divergence between them would silently train on off-policy tokens. + + Raises ``ValueError`` for unparseable or multimodal messages. The + tokenizer-heavy part runs in a thread: the HF fast tokenizer releases the + GIL, so a 30k-token render stops serializing the caller's event loop. + """ + from tensorrt_llm.serve.chat_utils import parse_chat_messages_coroutines + + conversation, mm_coroutine, *_ = parse_chat_messages_coroutines( + messages, model_config + ) + mm_data, mm_embeddings = await mm_coroutine + if mm_data is not None or mm_embeddings is not None: + raise ValueError( + "NeMo-RL's TRT-LLM HTTP adapter does not support multimodal chat inputs" + ) + + def _sync() -> list[int]: + # Full retokenization avoids accumulating generation token IDs twice. + prompt_token_ids = _build_prompt_token_ids( + conversation, + tokenizer, + tools=tools, + default_template_kwargs=template_kwargs, + ) + # Empty required_prefix_ids on turn one returns the template unchanged. + required_prefix_ids, template_prefix_ids = _compute_splice_inputs( + messages, + conversation, + tokenizer, + tools, + template_kwargs, + ) + return replace_prefix_tokens( + tokenizer=tokenizer, + model_prefix_token_ids=required_prefix_ids, + template_prefix_token_ids=template_prefix_ids, + template_token_ids=prompt_token_ids, + ) + + return await asyncio.to_thread(_sync) + + +# Sampling rate for shadow-validating frontend-supplied prompt ids on the +# context leg: the ctx adapter recomputes the ids from the messages it also +# received and logs a mismatch. Cheap insurance while frontend_tokenize is +# young; 0 disables. +_TOKENIZE_SHADOW_RATE = float( + os.environ.get("NRL_TRTLLM_TOKENIZE_SHADOW_RATE", "0") or 0 +) + + +def _context_leg_response( + model_name: str, + prompt_token_ids: list[int], + gen: Any, + disagg_params: Any, +) -> Any: + """Reply to a ``context_only`` request. + + Prefill materialised KV and at most one token; the disagg server only reads + the handshake back off this response (plus the prompt token ids, so the + generation server need not re-tokenize). Nothing here is user-visible, so + the reasoning/tool/stop-token post-processing is skipped entirely. + """ + from fastapi.responses import JSONResponse + from tensorrt_llm.serve.openai_protocol import to_disaggregated_params + + ctx_out = getattr(gen, "disaggregated_params", None) + if ctx_out is None: + raise RuntimeError( + "context leg returned no disaggregated_params; the engine is most " + "likely missing cache_transceiver_config" + ) + + response: dict[str, Any] = { + "id": f"chatcmpl-{uuid.uuid4().hex[:12]}", + "object": "chat.completion", + "created": int(time.time()), + "model": model_name, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": None}, + "finish_reason": gen.finish_reason, + "disaggregated_params": to_disaggregated_params(ctx_out).model_dump(), + } + ], + "usage": { + "prompt_tokens": len(prompt_token_ids), + "completion_tokens": 0, + "total_tokens": len(prompt_token_ids), + }, + } + + # The orchestrator asks for the base64 int32 buffer when it wants to relay a + # string instead of materialising the int list on its event loop. + if getattr(disagg_params, "return_prompt_token_ids_b64", False): + import base64 + + import numpy as np + + response["prompt_token_ids_b64"] = base64.b64encode( + np.asarray(prompt_token_ids, dtype=np.int32).tobytes() + ).decode("ascii") + else: + response["prompt_token_ids"] = prompt_token_ids + + return JSONResponse(content=response) + + def _build_reasoning_parser(name: str, chat_template_kwargs: dict[str, Any]) -> Any: from tensorrt_llm.llmapi.reasoning_parser import ReasoningParserFactory @@ -112,8 +271,6 @@ def create_app( _tool_parser_instance = _build_tool_parser(_tool_parser_name) _parse_tool_calls = _make_parse_tool_calls(_tool_parser_instance) - from tensorrt_llm.serve.chat_utils import parse_chat_messages_coroutines - model_config = getattr(llm, "_hf_model_config", None) if model_config is None: raise RuntimeError( @@ -157,6 +314,44 @@ async def chat_completions(request: Request): tools: list[dict] | None = body.get("tools") logprobs_requested = body.get("logprobs", False) + # Under PD disaggregation a replica's OpenAIDisaggServer drives this + # endpoint twice per request -- once context_only, once generation_only + # -- carrying the handshake between the two. The wire model differs from + # the engine one (opaque_state is bytes in the engine, base64 on the + # wire), so use TRT-LLM's own converter rather than reproducing it. + # Conversation identity for rank-affine ADP routing: canonical body + # conversation_params, else the id the disagg service stamps onto + # disaggregated_params for its ctx/gen legs. None = no affinity. + # + # The canonical field is not populated yet: Gym is the only component + # that knows the rollout identity, and no released Gym sends it (it + # appears nowhere in the tree, including upstream main). It arrives with + # https://github.com/NVIDIA-NeMo/Gym/pull/3582 ("forward gym session id + # as backend conversation id"), still open at time of writing, after + # which the Gym submodule pin has to be bumped for this branch to see an + # id at all. + # + # Until then a turn lands on an effectively random attention-DP rank and + # the context engine re-prefills most of its history: measured at + # 2P-DEP8 / conc 512, 17-26% of turns found their prefix on the serving + # rank (about 1/DEP) versus 96-97% once the id flows. Nothing fails -- + # disaggregation just gives up most of its benefit -- so treat a run + # with no conversation id as unmeasured rather than as a baseline. + _conv_id = (body.get("conversation_params") or {}).get("conversation_id") or ( + body.get("disaggregated_params") or {} + ).get("conversation_id") + disagg_params = None + if body.get("disaggregated_params") is not None: + from tensorrt_llm.serve.openai_protocol import ( + DisaggregatedParams as WireDisaggregatedParams, + ) + from tensorrt_llm.serve.openai_protocol import to_llm_disaggregated_params + + disagg_params = to_llm_disaggregated_params( + WireDisaggregatedParams(**body["disaggregated_params"]) + ) + is_context_leg = getattr(disagg_params, "request_type", None) == "context_only" + # The NeMo-RL generation config, not the request, is the source of truth # for sampling params. for key in ("temperature", "top_p", "top_k"): @@ -176,47 +371,84 @@ async def chat_completions(request: Request): else None ) - try: - conversation, mm_coroutine, _ = parse_chat_messages_coroutines( - messages, model_config - ) - mm_data, mm_embeddings = await mm_coroutine - except ValueError as e: - return JSONResponse(status_code=400, content={"error": str(e)}) - - # This token-only adapter does not support multimodal inputs. - if mm_data is not None or mm_embeddings is not None: - return JSONResponse( - status_code=400, - content={ - "error": "NeMo-RL's TRT-LLM HTTP adapter does not support " - "multimodal chat inputs" - }, - ) - - # Full retokenization avoids accumulating generation token IDs twice. - prompt_token_ids = _build_prompt_token_ids( - conversation, - tokenizer, - tools=tools, - default_template_kwargs=effective_template_kwargs, - ) - - # Empty required_prefix_ids on turn one returns the template unchanged. - required_prefix_ids, template_prefix_ids = _compute_splice_inputs( - messages, - conversation, - tokenizer, - tools, - effective_template_kwargs, - ) - - adj_prompt = replace_prefix_tokens( - tokenizer=tokenizer, - model_prefix_token_ids=required_prefix_ids, - template_prefix_token_ids=template_prefix_ids, - template_token_ids=prompt_token_ids, - ) + # On the generation leg the disagg server hands over the exact token ids + # the context engine built KV for (openai_disagg_service._get_gen_request). + # Rebuilding them from `messages` could yield a different sequence, which + # would decode against mismatched KV -- and silently. Prefer what it sent, + # and skip the chat-template work entirely: nothing downstream of this + # block reads `conversation`, and under gen_strip_message_history the + # messages are not even complete. (Before this short-circuit the + # generation leg rendered two full 30k-token templates per request and + # then discarded them.) + supplied = body.get("prompt_token_ids") + if supplied is None and body.get("prompt_token_ids_b64"): + # Same int32 buffer encoding openai_server.py uses on this hop. + import base64 + + import numpy as np + + supplied = np.frombuffer( + base64.b64decode(body["prompt_token_ids_b64"]), dtype=np.int32 + ).tolist() + + if supplied: + adj_prompt = list(supplied) + # Shadow validation for frontend-tokenized ids (context leg only: + # the generation leg's ids legitimately extend past the messages). + if ( + is_context_leg + and messages + and _TOKENIZE_SHADOW_RATE > 0 + and random.random() < _TOKENIZE_SHADOW_RATE + ): + try: + recomputed = await build_spliced_prompt_ids( + messages, + tools, + tokenizer, + model_config, + effective_template_kwargs, + ) + if recomputed != adj_prompt: + diff = next( + ( + i + for i, (a, b) in enumerate(zip(adj_prompt, recomputed)) + if a != b + ), + min(len(adj_prompt), len(recomputed)), + ) + lo, hi = max(0, diff - 8), diff + 8 + logger.error( + "frontend-tokenize SHADOW MISMATCH: supplied %d ids, " + "recomputed %d, first divergence at %d; " + "supplied[%d:%d]=%s (%r) vs recomputed=%s (%r); " + "template_kwargs=%r first_msg_role=%r", + len(adj_prompt), + len(recomputed), + diff, + lo, + hi, + adj_prompt[lo:hi], + tokenizer.decode(adj_prompt[lo:hi]), + recomputed[lo:hi], + tokenizer.decode(recomputed[lo:hi]), + effective_template_kwargs, + (messages[0] or {}).get("role") if messages else None, + ) + except Exception as e: # noqa: BLE001 - shadow must never fail serving + logger.error("frontend-tokenize shadow recompute failed: %s", e) + else: + try: + adj_prompt = await build_spliced_prompt_ids( + messages, + tools, + tokenizer, + model_config, + effective_template_kwargs, + ) + except ValueError as e: + return JSONResponse(status_code=400, content={"error": str(e)}) max_tokens_requested = ( body.get("max_tokens") or body.get("max_completion_tokens") or max_seq_len @@ -245,9 +477,16 @@ async def chat_completions(request: Request): ) try: + _conv_params = None + if _conv_id: + from tensorrt_llm.conversation_params import ConversationParams + + _conv_params = ConversationParams(conversation_id=str(_conv_id)) output = await llm.generate_async( {"prompt_token_ids": adj_prompt}, sampling_params=sampling, + disaggregated_params=disagg_params, + conversation_params=_conv_params, ) except RequestError as e: err = str(e) @@ -259,6 +498,14 @@ async def chat_completions(request: Request): raise gen = output.outputs[0] + + if is_context_leg: + # Prefill produced KV and at most one token. Everything downstream + # -- reasoning parsing, tool parsing, stop-token trimming -- is for + # the completed generation, so skip it and hand the disagg server + # just what it needs to build the generation leg. + return _context_leg_response(model_name, adj_prompt, gen, disagg_params) + gen_token_ids = list(gen.token_ids) gen_logprobs: list[float] = [] @@ -307,9 +554,6 @@ async def chat_completions(request: Request): "content": content_text or None, "reasoning_content": reasoning_content, "tool_calls": parsed_tool_calls, - "prompt_token_ids": adj_prompt, - "generation_token_ids": gen_token_ids, - "generation_log_probs": gen_logprobs, } finish_reason = "tool_calls" else: @@ -317,19 +561,40 @@ async def chat_completions(request: Request): "role": "assistant", "content": answer_text, "reasoning_content": reasoning_content, - "prompt_token_ids": adj_prompt, - "generation_token_ids": gen_token_ids, - "generation_log_probs": gen_logprobs, } + # NeMo-Gym reads the rollout fields off the *message* + # (nemo_rl/environments/nemo_gym.py: a message without + # generation_token_ids is skipped outright, so a miss loses the whole + # turn's training data silently). Aggregated serving answers Gym + # directly, so attach them here. + # + # Not under disaggregation: there the reply is re-validated by the + # disagg server against ChatMessage, which is extra="forbid" and would + # 400 on these. They ride the declared fields instead + # (choices[].token_ids, prompt_token_ids, logprobs) and the disagg + # server's outbound adaptor re-attaches them to the message before Gym + # ever sees it -- see trtllm_disagg_server._attach_rollout_fields. + if disagg_params is None: + msg_dict["prompt_token_ids"] = adj_prompt + msg_dict["generation_token_ids"] = gen_token_ids + msg_dict["generation_log_probs"] = gen_logprobs + response: dict[str, Any] = { "id": f"chatcmpl-{uuid.uuid4().hex[:12]}", "object": "chat.completion", "created": int(time.time()), "model": model_name, "choices": [ - {"index": 0, "message": msg_dict, "finish_reason": finish_reason} + { + "index": 0, + "message": msg_dict, + "finish_reason": finish_reason, + } ], + # Declared on ChatCompletionResponse precisely so a generation + # server need not re-tokenize the prompt. + "prompt_token_ids": adj_prompt, "usage": { "prompt_tokens": len(adj_prompt), "completion_tokens": len(gen_token_ids), @@ -337,11 +602,37 @@ async def chat_completions(request: Request): }, } + # Generated token ids get their own declared field only on a TRT-LLM + # build whose ChatCompletionResponseChoice carries one + # (CompletionResponseChoice already does). Emitting it unconditionally + # would 400 at the disagg server's extra="forbid" re-validation on every + # other build -- which is precisely where the token_id:N logprobs + # fallback below is supposed to take over, so the write has to be gated + # the same way the read is. + if _choice_declares_token_ids(): + response["choices"][0]["token_ids"] = gen_token_ids + if logprobs_requested and gen_logprobs: + # `token` carries the id rather than the decoded text when asked. + # ChatCompletionResponseChoice has no token-id field upstream yet, so + # this declared string field is how the ids survive the disagg + # server's strict re-validation. Same encoding vLLM uses, which is + # what NeMo-Gym already parses. + # Under disaggregation the flag never survives the frontend's + # gym-field strip, yet ids are exactly what rides this channel + # (the disagg adaptor parses token_id:N back out) -- and decoding + # each generated token individually is per-token tokenizer work + # nobody reads. Force the id encoding on any disagg request. + as_ids = ( + bool(body.get("return_tokens_as_token_ids")) + or disagg_params is not None + ) response["choices"][0]["logprobs"] = { "content": [ { - "token": tokenizer.decode([tid]), + "token": ( + f"token_id:{tid}" if as_ids else tokenizer.decode([tid]) + ), "logprob": lp, "bytes": None, "top_logprobs": [], diff --git a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py index 32839ada63a..b964f58d5a3 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py +++ b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py @@ -44,6 +44,29 @@ from nemo_rl.models.generation.trtllm.config import TrtllmConfig +def _tag_nsys_output(output_spec: str, tag: str) -> str: + """Append *tag* to the filename in an nsys ``-o`` spec. + + Operates on the filename rather than substituting a known worker name, + because the spec may be the built-in default or anything the user put in + ``NRL_NSYS_EXTRA_OPTIONS`` -- including a directory, which must be left + where it is so the reports still land where the user asked. The value + arrives quoted from ``get_nsight_config_if_pattern_matches``; a + user-supplied one may not be, so the quoting is preserved either way. + + Appended rather than prefixed so the documented report names stay + prefix-matchable: ``trtllm_async_generation_worker_*`` keeps finding every + report (docs/nsys-profiling.md), disaggregated or not. + """ + quote = "" + if len(output_spec) >= 2 and output_spec[0] == output_spec[-1]: + if output_spec[0] in "'\"": + quote, output_spec = output_spec[0], output_spec[1:-1] + + head, sep, name = output_spec.rpartition("/") + return f"{quote}{head}{sep}{name}_{tag}{quote}" + + class TrtllmAsyncGenerationWorkerImpl: """Plain (non-actor) implementation of the async TRT-LLM generation worker. @@ -79,6 +102,12 @@ def configure_worker( # parent placement group via get_current_placement_group() and hands both # to TRT-LLM as ray_placement_config (instead of TRTLLM_RAY_BUNDLE_INDICES). init_kwargs["bundle_indices"] = bundle_indices[1] + # The placement-group index completes the engine's identity, which + # the worker uses to find its own entry in trtllm_cfg's + # _engine_overrides map (see TrtllmGeneration._engine_key). Local + # bundle indices alone are ambiguous across per-node PGs, which + # each restart numbering at 0. + init_kwargs["bundle_pg_idx"] = bundle_indices[0] init_kwargs["fraction_of_gpus"] = num_gpus @@ -87,26 +116,59 @@ def configure_worker( def __repr__(self) -> str: return "TrtllmAsyncGenerationWorker" + def _engine_overrides(self) -> dict[str, Any]: + """This engine's entry in the driver-built ``_engine_overrides`` map. + + Keyed by ``(placement group index, local bundle indices)`` -- the same + tuple ``TrtllmGeneration`` used to create this worker, so the two sides + agree without per-worker init kwargs. Returns ``{}`` for uniform runs, + which carry no map at all. + """ + overrides = self.cfg["trtllm_cfg"].get("_engine_overrides") + if not overrides or self._bundle_indices is None: + return {} + key = f"{self._bundle_pg_idx}:" + ",".join(str(i) for i in self._bundle_indices) + entry = overrides.get(key) + if entry is None: + raise RuntimeError( + f"No engine override entry for {key!r}. Known keys: " + f"{sorted(overrides)}. TrtllmGeneration._engine_key() and this " + f"lookup must stay in sync." + ) + return entry + def __init__( self, config: TrtllmConfig, bundle_indices: Optional[list[int]] = None, + bundle_pg_idx: Optional[int] = None, fraction_of_gpus: float = 1.0, seed: Optional[int] = None, ) -> None: self.cfg = config + self._bundle_pg_idx = bundle_pg_idx # Allow gen side to use a quantized checkpoint self.model_name = ( self.cfg.get("trtllm_cfg", {}).get("model_name") or self.cfg["model_name"] ) self.is_model_owner = bundle_indices is not None self._bundle_indices = bundle_indices + # This engine's effective config: the shared trtllm_cfg with this + # engine's role overrides merged over it. Everything that describes the + # engine -- constructor kwargs, the HTTP server's limits -- must read + # this, never the raw trtllm_cfg, or a role's overrides apply to some + # settings and not others. Equal to trtllm_cfg without disaggregation. + self.engine_cfg: dict[str, Any] = { + **self.cfg["trtllm_cfg"], + **self._engine_overrides(), + } self._fraction_of_gpus = fraction_of_gpus self._seed = seed self.llm = None self.TrtSamplingParams = None self._http_thread = None self._http_base_url: Optional[str] = None + self._http_addr: Optional[tuple[str, int]] = None self._http_server = None if not self.is_model_owner: @@ -126,8 +188,13 @@ def __init__( self.TrtSamplingParams = TrtSamplingParams - trtllm_cfg = self.cfg["trtllm_cfg"] - tp_size = trtllm_cfg["tensor_parallel_size"] + engine_cfg = self.engine_cfg + # This engine's TP is the number of bundles TrtllmGeneration tied to it, + # not the config value: under PD disaggregation with asymmetric TP the + # context and generation engines are different widths, and the config + # holds only the default. Identical to engine_cfg["tensor_parallel_size"] + # whenever the layout is uniform. + tp_size = len(self._bundle_indices) self._colocated = self.cfg["colocated"]["enabled"] os.environ.pop("CUDA_VISIBLE_DEVICES", None) @@ -157,15 +224,15 @@ def __init__( model=self.model_name, backend="pytorch", tensor_parallel_size=tp_size, - dtype=trtllm_cfg["precision"], - max_seq_len=trtllm_cfg["max_model_len"], - max_batch_size=trtllm_cfg["max_batch_size"], - max_num_tokens=trtllm_cfg["max_num_tokens"], + dtype=engine_cfg["precision"], + max_seq_len=engine_cfg["max_model_len"], + max_batch_size=engine_cfg["max_batch_size"], + max_num_tokens=engine_cfg["max_num_tokens"], # vLLM accepts prompts up to max_model_len (no separate input cap; it clamps output so # input+output <= max_model_len). TRT-LLM defaults max_input_len=1024, which rejects long # SWE-agent prompts before any tokens generate -> NeMo Gym sees "no generation data". # Match the input cap to the context window so it isn't the bottleneck. - max_input_len=trtllm_cfg["max_model_len"], + max_input_len=engine_cfg["max_model_len"], orchestrator_type="ray", ray_worker_extension_cls="nemo_rl.models.generation.trtllm.trtllm_backend.NcclExtension", placement_groups=placement_groups_list, @@ -176,7 +243,7 @@ def __init__( ), cuda_graph_config=CudaGraphConfig( enable_padding=True, - max_batch_size=trtllm_cfg["max_batch_size"], + max_batch_size=engine_cfg["max_batch_size"], ), ) @@ -186,19 +253,59 @@ def __init__( # AsyncLLM kwargs (which are validated against LlmArgs.model_fields and # reject unknown keys) — pass them via kv_cache_config instead. The rest # of trtllm_kwargs is spread as top-level AsyncLLM kwargs below. + # A role may override kv_cache_config too: prefill and decode engines + # usually want different KV pool sizes. extra_trtllm_kwargs = dict(self.cfg.get("trtllm_kwargs") or {}) + extra_trtllm_kwargs.update(engine_cfg.get("trtllm_kwargs") or {}) kv_cache_kwargs = dict(extra_trtllm_kwargs.pop("kv_cache_config", None) or {}) - # gpu_memory_utilization is a dedicated trtllm_cfg knob for + # gpu_memory_utilization is a dedicated engine_cfg knob for # free_gpu_memory_fraction; an explicit kv_cache_config value wins. - gpu_mem_util = trtllm_cfg.get("gpu_memory_utilization") + gpu_mem_util = engine_cfg.get("gpu_memory_utilization") if gpu_mem_util is not None: kv_cache_kwargs.setdefault("free_gpu_memory_fraction", gpu_mem_util) if kv_cache_kwargs: llm_kwargs["kv_cache_config"] = KvCacheConfig(**kv_cache_kwargs) - moe_tp = trtllm_cfg.get("moe_tensor_parallel_size") - moe_ep = trtllm_cfg.get("moe_expert_parallel_size") + # PD disaggregation: every engine gets a cache transceiver. The + # context/generation role is per *request* in TRT-LLM + # (DisaggregatedParams.request_type), not per engine, so engines are + # symmetric here; which pool an engine lands in is decided by the + # replica's disagg server from its address alone. + disagg_cfg = engine_cfg.get("disaggregation") or {} + if disagg_cfg.get("enabled"): + from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig + + transceiver_kwargs: dict[str, Any] = { + "backend": disagg_cfg["cache_transceiver_backend"], + } + if disagg_cfg.get("max_tokens_in_buffer") is not None: + transceiver_kwargs["max_tokens_in_buffer"] = disagg_cfg[ + "max_tokens_in_buffer" + ] + # Only forward when set, so TRT-LLM keeps its own "auto" default + # otherwise. Worth forwarding at all because "auto" resolves to the + # C++ transceiver whenever it cannot confirm the model's preference, + # and a hybrid Mamba model under disaggregation needs the Python + # (v2) transceiver to hand its recurrent state over. + if disagg_cfg.get("cache_transceiver_runtime") is not None: + transceiver_kwargs["transceiver_runtime"] = disagg_cfg[ + "cache_transceiver_runtime" + ] + if disagg_cfg.get("kv_cache_bounce_size_mb") is not None: + transceiver_kwargs["kv_cache_bounce_size_mb"] = disagg_cfg[ + "kv_cache_bounce_size_mb" + ] + if disagg_cfg.get("kv_transfer_timeout_ms") is not None: + transceiver_kwargs["kv_transfer_timeout_ms"] = disagg_cfg[ + "kv_transfer_timeout_ms" + ] + llm_kwargs["cache_transceiver_config"] = CacheTransceiverConfig( + **transceiver_kwargs + ) + + moe_tp = engine_cfg.get("moe_tensor_parallel_size") + moe_ep = engine_cfg.get("moe_expert_parallel_size") if moe_tp is not None: llm_kwargs["moe_tensor_parallel_size"] = moe_tp if moe_ep is not None: @@ -231,6 +338,21 @@ def __init__( "trtllm_async_generation_worker" ).get("nsight") if _nsight and "ray_worker_nsight_options" not in llm_kwargs: + _disagg_role = self.engine_cfg.get("_disagg_role") + if _disagg_role: + # Under disaggregation every engine reports the same worker + # name, so the report name distinguishes them only by the %p + # pid -- which means matching a report to the context or + # generation side is a search through the driver log. Tag the + # filename with the role and its ordinal (a layout can run + # several engines of one role); %p still separates the ranks + # within an engine. Aggregated runs have no role and are left + # alone. + _nsight = dict(_nsight) + _nsight["o"] = _tag_nsys_output( + _nsight["o"], + f"{_disagg_role}{self.engine_cfg['_disagg_role_ordinal']}", + ) llm_kwargs["ray_worker_nsight_options"] = _nsight # Defer __await__ (which fires setup_async) to post_init_async so @@ -254,6 +376,9 @@ async def post_init_async(self) -> None: print("[TrtllmAsyncWorker] AsyncLLM ready", flush=True) if self.cfg["trtllm_cfg"].get("expose_http_server"): + # Every engine serves HTTP, context and generation alike: under + # disaggregation an address is the only way the disagg server can + # reach one. self.start_http_server() def shutdown(self) -> bool: @@ -291,19 +416,28 @@ def start_http_server(self, port: int = 0) -> str: tokenizer=tokenizer, model_name=self.model_name, port=port, - max_seq_len=self.cfg["trtllm_cfg"]["max_model_len"], + # engine_cfg, not trtllm_cfg: the server must not accept requests + # longer than the engine it fronts was built for, and a role may + # have overridden that length. + max_seq_len=self.engine_cfg["max_model_len"], sampling_config={ "temperature": self.cfg["temperature"], "top_p": self.cfg["top_p"], "top_k": self.cfg["top_k"], }, stop_token_ids=list(self.cfg.get("stop_token_ids") or []), - default_chat_template_kwargs=self.cfg["trtllm_cfg"].get( + default_chat_template_kwargs=self.engine_cfg.get( "default_chat_template_kwargs" ), - tool_parser=self.cfg["trtllm_cfg"].get("tool_parser"), - reasoning_parser=self.cfg["trtllm_cfg"].get("reasoning_parser"), + tool_parser=self.engine_cfg.get("tool_parser"), + reasoning_parser=self.engine_cfg.get("reasoning_parser"), ) + # host:port of the URL just returned; the disagg server is configured + # with addresses, not URLs. + _hostport = self._http_base_url.rsplit("/v1", 1)[0].rsplit("//", 1)[1] + _host, _port = _hostport.rsplit(":", 1) + self._http_addr = (_host, int(_port)) + print( f"[TrtllmAsyncWorker] HTTP server started: {self._http_base_url}", flush=True, @@ -316,10 +450,28 @@ def stop_http_server(self) -> None: self._http_server = None self._http_thread = None self._http_base_url = None + self._http_addr = None async def report_dp_openai_server_base_url(self) -> Optional[str]: return self._http_base_url + async def report_http_addr(self) -> Optional[dict[str, Any]]: + """This engine's ``{host, port, node_id}``. + + Under disaggregation an address is the only thing the replica's disagg + server is given about an engine, so host/port is what feeds its context + and generation pools. The node id lets the driver put that server's + actor beside its engines. + """ + if self._http_addr is None: + return None + host, port = self._http_addr + return { + "host": host, + "port": port, + "node_id": ray.get_runtime_context().get_node_id(), + } + # ------------------------------------------------------------------ # # Collective RPC / refit # ------------------------------------------------------------------ # diff --git a/pyrefly.toml b/pyrefly.toml index e4ffa4a3673..1ca0edd1c07 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -231,6 +231,7 @@ project-includes = [ "nemo_rl/models/generation/sglang/utils/ray_utils.py", "nemo_rl/models/generation/trtllm/__init__.py", "nemo_rl/models/generation/trtllm/config.py", + "nemo_rl/models/generation/trtllm/trtllm_disagg_server.py", "nemo_rl/models/generation/trtllm/trtllm_generation.py", "nemo_rl/models/generation/trtllm/trtllm_http_server.py", "nemo_rl/models/generation/vllm/__init__.py", diff --git a/ray.sub b/ray.sub index 7caaeb45497..378ddf11a35 100644 --- a/ray.sub +++ b/ray.sub @@ -264,6 +264,14 @@ COMMON_SRUN_ARGS+=" --container-workdir=$SLURM_SUBMIT_DIR" # Pass partition/account explicitly for overlapping srun calls. COMMON_SRUN_ARGS+=" -p $SLURM_JOB_PARTITION" COMMON_SRUN_ARGS+=" -A $SLURM_JOB_ACCOUNT" +# Opt-in: force host values of these comma-separated variables over the image's +# baked ENV (pyxis lets the image win by default). Needed e.g. for UCX_NET_DEVICES=all +# so NIXL KV transfer uses the RDMA HCAs instead of the image's management NIC. +if [[ -n "${NRL_CONTAINER_ENV_FORWARD:-}" ]]; then + for _v in ${NRL_CONTAINER_ENV_FORWARD//,/ }; do + COMMON_SRUN_ARGS+=" --container-env=${_v}" + done +fi # Number of CPUs per worker node. If the user did not set CPUS_PER_WORKER # explicitly, auto-detect it from SLURM so we claim every CPU the node actually # has (instead of the old GPUS_PER_NODE * 16 heuristic, which under-claims on diff --git a/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py b/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py new file mode 100644 index 00000000000..afcd2a0aa9a --- /dev/null +++ b/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py @@ -0,0 +1,282 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The disagg frontend's request-adapting middleware. + +TRT-LLM's OpenAIDisaggServer validates bodies against ``extra="forbid"`` +models, so any vLLM-only request field must be stripped before it reaches the +route. ``required_prefix_token_ids`` (the vLLM prefix-override extension the +MLPerf warmup sends) turned every warmup request into a 400 until it was added +to the strip list; the TRT-LLM engine adapter ignores that field anyway. +""" + +import asyncio +import json + +import pytest + +from nemo_rl.models.generation.trtllm.trtllm_disagg_server import ( + _GYM_ONLY_REQUEST_FIELDS, + _DropGymOnlyRequestFields, + attach_rollout_fields, + generation_token_ids, +) + +pytestmark = pytest.mark.trtllm + + +class _CaptureApp: + """Downstream ASGI app recording the scope and body it is handed.""" + + def __init__(self): + self.scope = None + self.body = b"" + + async def __call__(self, scope, receive, send): + self.scope = scope + chunks = [] + while True: + message = await receive() + if message["type"] != "http.request": + break + chunks.append(message.get("body", b"")) + if not message.get("more_body", False): + break + self.body = b"".join(chunks) + + +def _scope(path: str, body: bytes) -> dict: + return { + "type": "http", + "path": path, + "headers": [ + (b"content-type", b"application/json"), + (b"content-length", str(len(body)).encode()), + ], + } + + +def _drive(middleware, scope, body: bytes) -> None: + async def receive(): + return {"type": "http.request", "body": body, "more_body": False} + + async def send(_message): + return None + + asyncio.run(middleware(scope, receive, send)) + + +def _warmup_payload() -> dict: + return { + "model": "qwen", + "messages": [{"role": "user", "content": "MLPerf warmup 0-0"}], + "required_prefix_token_ids": [11, 22, 33], + "return_token_ids": True, + "max_tokens": 8, + "temperature": 1.0, + } + + +def test_required_prefix_token_ids_is_a_stripped_field(): + assert "required_prefix_token_ids" in _GYM_ONLY_REQUEST_FIELDS + + +def test_vllm_only_fields_are_stripped_and_content_length_follows(): + payload = _warmup_payload() + body = json.dumps(payload).encode() + app = _CaptureApp() + + _drive(_DropGymOnlyRequestFields(app), _scope("/v1/chat/completions", body), body) + + got = json.loads(app.body) + assert "required_prefix_token_ids" not in got + assert "return_token_ids" not in got + # Everything the engine needs survives untouched. + assert got["messages"] == payload["messages"] + assert got["max_tokens"] == 8 and got["temperature"] == 1.0 + headers = dict(app.scope["headers"]) + assert int(headers[b"content-length"]) == len(app.body) + + +def test_frontend_tokenizer_never_sees_the_stripped_fields(): + seen = [] + + async def tokenize_fn(payload): + seen.append(dict(payload)) + return "QUJD" # base64("ABC") + + payload = _warmup_payload() + body = json.dumps(payload).encode() + app = _CaptureApp() + + _drive( + _DropGymOnlyRequestFields(app, tokenize_fn=tokenize_fn), + _scope("/v1/chat/completions", body), + body, + ) + + assert len(seen) == 1 + assert "required_prefix_token_ids" not in seen[0] + assert json.loads(app.body)["prompt_token_ids_b64"] == "QUJD" + + +def test_non_adapted_paths_pass_through_untouched(): + body = b'{"required_prefix_token_ids": [1]}' + scope = _scope("/v1/models", body) + app = _CaptureApp() + + _drive(_DropGymOnlyRequestFields(app), scope, body) + + assert app.body == body + assert app.scope is scope + + +def test_a_failing_frontend_tokenizer_falls_back_instead_of_failing_the_request(): + async def tokenize_fn(payload): + raise RuntimeError("tokenizer died") + + payload = _warmup_payload() + body = json.dumps(payload).encode() + app = _CaptureApp() + + _drive( + _DropGymOnlyRequestFields(app, tokenize_fn=tokenize_fn), + _scope("/v1/chat/completions", body), + body, + ) + + # No b64 ids, so the ctx adapter tokenizes from the messages itself -- + # degraded, not a 500. + got = json.loads(app.body) + assert "prompt_token_ids_b64" not in got + assert "required_prefix_token_ids" not in got + assert got["messages"] == payload["messages"] + + +def test_engine_facing_legs_are_not_retokenized(): + calls = [] + + async def tokenize_fn(payload): + calls.append(payload) + return "QUJD" + + payload = _warmup_payload() + payload["disaggregated_params"] = {"request_type": "generation_only"} + body = json.dumps(payload).encode() + app = _CaptureApp() + + _drive( + _DropGymOnlyRequestFields(app, tokenize_fn=tokenize_fn), + _scope("/v1/chat/completions", body), + body, + ) + + # The generation leg's history may already be stripped, so re-deriving the + # prompt ids from it would contradict the KV the context leg transferred. + assert calls == [] + assert "prompt_token_ids_b64" not in json.loads(app.body) + + +def test_a_chunked_body_is_reassembled_before_the_fields_are_stripped(): + payload = _warmup_payload() + body = json.dumps(payload).encode() + half = len(body) // 2 + app = _CaptureApp() + middleware = _DropGymOnlyRequestFields(app) + + parts = iter( + [ + {"type": "http.request", "body": body[:half], "more_body": True}, + {"type": "http.request", "body": body[half:], "more_body": False}, + ] + ) + + async def receive(): + return next(parts) + + async def send(_message): + return None + + asyncio.run(middleware(_scope("/v1/chat/completions", body), receive, send)) + + got = json.loads(app.body) + assert "required_prefix_token_ids" not in got + assert got["messages"] == payload["messages"] + + +def test_a_non_json_body_is_forwarded_verbatim(): + body = b"not json at all" + app = _CaptureApp() + + _drive(_DropGymOnlyRequestFields(app), _scope("/v1/completions", body), body) + + assert app.body == body + + +# -------------------------------------------------------------------------- # +# Outbound: the rollout fields NeMo-Gym reads off choices[].message +# -------------------------------------------------------------------------- # + + +def _disagg_response(*, token_ids=None, logprobs_content=None) -> dict: + """The shape trtllm_http_server emits on the disagg (non-aggregated) path.""" + choice: dict = { + "index": 0, + "message": {"role": "assistant", "content": "hi", "reasoning_content": None}, + "finish_reason": "stop", + } + if token_ids is not None: + choice["token_ids"] = token_ids + if logprobs_content is not None: + choice["logprobs"] = {"content": logprobs_content} + return {"choices": [choice], "prompt_token_ids": [7, 8, 9]} + + +def test_rollout_fields_move_onto_the_message_from_the_declared_token_ids_field(): + payload = attach_rollout_fields( + _disagg_response( + token_ids=[4, 5], + logprobs_content=[ + {"token": "token_id:4", "logprob": -0.5}, + {"token": "token_id:5", "logprob": -1.5}, + ], + ) + ) + + message = payload["choices"][0]["message"] + assert message["prompt_token_ids"] == [7, 8, 9] + assert message["generation_token_ids"] == [4, 5] + assert message["generation_log_probs"] == [-0.5, -1.5] + + +def test_generation_token_ids_fall_back_to_the_logprobs_encoding(): + # The build in use has no ChatCompletionResponseChoice.token_ids, so the ids + # ride the declared logprobs strings instead. + payload = attach_rollout_fields( + _disagg_response( + logprobs_content=[ + {"token": "token_id:11", "logprob": -0.1}, + {"token": "token_id:22", "logprob": -0.2}, + ] + ) + ) + + assert payload["choices"][0]["message"]["generation_token_ids"] == [11, 22] + + +def test_decoded_tokens_are_not_mistaken_for_ids(): + assert generation_token_ids({"logprobs": {"content": [{"token": "hello"}]}}) is None + + +def test_a_response_without_choices_is_returned_unchanged(): + assert attach_rollout_fields({"choices": []}) == {"choices": []} diff --git a/tests/unit/models/generation/trtllm/test_trtllm_generation.py b/tests/unit/models/generation/trtllm/test_trtllm_generation.py index 1123ebae58a..98ef57c02d9 100644 --- a/tests/unit/models/generation/trtllm/test_trtllm_generation.py +++ b/tests/unit/models/generation/trtllm/test_trtllm_generation.py @@ -18,9 +18,11 @@ import pytest import torch +from pydantic import ValidationError from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.models.generation.trtllm import trtllm_generation +from nemo_rl.models.generation.trtllm.config import resolve_trtllm_disagg_config from nemo_rl.models.generation.trtllm.trtllm_generation import TrtllmGeneration pytestmark = pytest.mark.trtllm @@ -336,3 +338,106 @@ def test_ipc_refit_and_missing_worker_group(): broken.worker_group.workers = [] with pytest.raises(RuntimeError, match="Worker group not initialised"): broken.update_weights_via_ipc_zmq() + + +# -------------------------------------------------------------------------- # +# Disaggregation config and engine planning +# -------------------------------------------------------------------------- # + + +def test_disagg_defaults_come_from_the_schema_not_the_call_sites(): + # Absent and present-but-empty must agree: every default lives on the model, + # so no consumer has to re-derive one per key. + for raw in (None, {}): + cfg = _config(disaggregation=raw) if raw is not None else _config() + disagg = resolve_trtllm_disagg_config(cfg) + assert disagg.enabled is False + assert disagg.num_context_engines == 1 + assert disagg.num_generation_engines == 1 + assert disagg.num_frontend_workers == 1 + assert disagg.frontend_base_port == 17300 + assert disagg.ctx_router == "conversation" + assert disagg.gen_router == "load_balancing" + assert disagg.cache_transceiver_backend == "DEFAULT" + assert disagg.cache_transceiver_runtime is None + assert disagg.ctx_trtllm_kwargs == {} + + +def test_a_plausible_but_wrong_router_is_rejected(): + # round_robin parses as a str but silently discards the prefix affinity the + # context engines depend on, so the schema has to reject it. + with pytest.raises(ValidationError): + resolve_trtllm_disagg_config( + _config(disaggregation={"enabled": True, "ctx_router": "round_robin"}) + ) + + +def test_unknown_disagg_keys_are_preserved_for_older_configs(): + disagg = resolve_trtllm_disagg_config( + _config(disaggregation={"enabled": True, "some_future_key": 3}) + ) + assert disagg.enabled is True + assert disagg.model_extra["some_future_key"] == 3 + + +@pytest.mark.parametrize( + "disaggregation, world_size, expected_roles, expected_tps, expected_replicas", + [ + # Disabled: a replica is an engine, so the count follows TP as before. + (None, 8, ["generation"] * 8, [1] * 8, 8), + # Enabled, 1:1 at TP1 -- the exemplar's defaults on one 8-GPU node. + ( + {"enabled": True}, + 8, + ["context", "generation"] * 4, + [1, 1] * 4, + 4, + ), + # Per-role TP and engine counts are independent, and same-role engines + # stay contiguous within a replica. + ( + { + "enabled": True, + "num_context_engines": 2, + "num_generation_engines": 1, + "ctx_trtllm_kwargs": {"tensor_parallel_size": 1}, + "gen_trtllm_kwargs": {"tensor_parallel_size": 2}, + }, + 8, + ["context", "context", "generation"] * 2, + [1, 1, 2] * 2, + 2, + ), + ], +) +def test_plan_engines_lays_replicas_out_contiguously( + disaggregation, world_size, expected_roles, expected_tps, expected_replicas +): + cfg = _config(**({"disaggregation": disaggregation} if disaggregation else {})) + generation = TrtllmGeneration.__new__(TrtllmGeneration) + generation.cfg = cfg + generation.tp_size = cfg["trtllm_cfg"]["tensor_parallel_size"] + generation._disagg = resolve_trtllm_disagg_config(cfg) + + roles, tps, num_replicas = generation._plan_engines(world_size) + + assert roles == expected_roles + assert tps == expected_tps + assert num_replicas == expected_replicas + + +def test_plan_engines_rejects_a_replica_width_that_does_not_tile_the_cluster(): + cfg = _config( + disaggregation={ + "enabled": True, + "ctx_trtllm_kwargs": {"tensor_parallel_size": 2}, + "gen_trtllm_kwargs": {"tensor_parallel_size": 1}, + } + ) + generation = TrtllmGeneration.__new__(TrtllmGeneration) + generation.cfg = cfg + generation.tp_size = cfg["trtllm_cfg"]["tensor_parallel_size"] + generation._disagg = resolve_trtllm_disagg_config(cfg) + + with pytest.raises(AssertionError, match="replica width 3 GPUs"): + generation._plan_engines(8) diff --git a/tests/unit/models/generation/trtllm/test_trtllm_http_server.py b/tests/unit/models/generation/trtllm/test_trtllm_http_server.py index fa5824d66ee..9ce2e29a9f2 100644 --- a/tests/unit/models/generation/trtllm/test_trtllm_http_server.py +++ b/tests/unit/models/generation/trtllm/test_trtllm_http_server.py @@ -26,6 +26,8 @@ _resolve_tool_parser_name, ) +pytestmark = pytest.mark.trtllm + class _FakeToolParser: def __init__(self, *, calls): diff --git a/tools/build-custom-trtllm.sh b/tools/build-custom-trtllm.sh index 380dd3bacbe..e541a3b354a 100755 --- a/tools/build-custom-trtllm.sh +++ b/tools/build-custom-trtllm.sh @@ -146,6 +146,13 @@ sed -i 's|nvidia-modelopt\[torch\]~=0\.37\.0|nvidia-modelopt[torch]>=0.44.0a0|' assert_patch_target requirements.txt 'setuptools<80' sed -i 's|^setuptools<80$|setuptools|' requirements.txt +# - drop PyNvVideoCodec. PyPI has no wheel in the pinned ~=2.1.0 range for +# aarch64/py3.13 (it jumps 2.0.5 -> 2.2.0), so build_wheel.py's +# `pip install -r requirements-dev.txt` aborts before cmake ever runs. +# Nothing in this build path decodes video. Not asserted: the pin only +# appeared after 1.3.0rc21, so older refs legitimately lack the line. +sed -i '/^PyNvVideoCodec/d' requirements.txt + # cutlass_kernels/CMakeLists.txt invokes `setup_library.py develop --user`, # which (a) requires a setup.py shim and (b) the `--user` flag is invalid # inside a venv. Rewrite the COMMAND to copy setup_library.py to setup.py @@ -181,6 +188,15 @@ echo "Building TensorRT-LLM wheel (arch=${ARCH}, jobs=${JOBS})..." show_ccache_status "before TRT-LLM build" ccache --zero-stats +# With NIXL enabled the build emits tensorrt_llm_transfer_agent_binding, which +# build_wheel.py imports to generate stubs. That module links torch but carries +# no rpath to it, so the import fails on "libc10.so: cannot open shared object +# file" and kills the build after the compile is already done. build_wheel.py +# copies the ambient environment for the stub step, so exporting torch's lib +# directory here is enough. +TORCH_LIB_DIR="$(python3 -c 'import os, torch; print(os.path.join(os.path.dirname(torch.__file__), "lib"))')" +export LD_LIBRARY_PATH="${TORCH_LIB_DIR}:${LD_LIBRARY_PATH:-}" + # Keep full output for failure diagnostics, but stream only 5% Ninja milestones. TRTLLM_BUILD_LOG=$(mktemp /tmp/trtllm-build.XXXXXX.log) set +x @@ -192,8 +208,24 @@ TRTLLM_BUILD_CMD=( --use_ccache --nvrtc_dynamic_linking --job_count "$JOBS" - -D "ENABLE_UCX=OFF" + # PD disaggregation's cache transceiver: ENABLE_UCX builds the UCX wrapper + # the engine dlopen()s, and it also gates find_package(NIXL) in + # cpp/CMakeLists.txt, so NIXL rides along with it. With UCX off, every + # non-MPI backend is compiled out and the engine aborts on the first KV + # transfer. + -D "ENABLE_UCX=ON" ) +# NIXL is what cache_transceiver_backend=DEFAULT resolves to, and it is +# installed by docker/Dockerfile. Missing means the image is wrong, so fail here +# rather than shipping a wheel whose default KV transport aborts at runtime. +NIXL_ROOT_DIR="${NIXL_ROOT_DIR:-/opt/nvidia/nvda_nixl}" +if [[ ! -f "${NIXL_ROOT_DIR}/include/nixl.h" ]]; then + echo "[ERROR] NIXL not found at ${NIXL_ROOT_DIR}. The image must install it" \ + "(see the NIXL_VERSION block in docker/Dockerfile), or point" \ + "NIXL_ROOT_DIR at an existing install." >&2 + exit 1 +fi +TRTLLM_BUILD_CMD+=(--nixl_root "${NIXL_ROOT_DIR}") set +e NINJA_STATUS='NINJA_PROGRESS:%f:%t:%p:' \ PYTHONUNBUFFERED=1 \