From f960aed1fabb2a169d1236d041a73ea08d7f9ec6 Mon Sep 17 00:00:00 2001 From: shuyixiong <219646547+shuyixiong@users.noreply.github.com> Date: Sun, 9 Aug 2026 19:38:26 -0700 Subject: [PATCH 01/12] feat(trtllm): TRT-LLM prefill/decode disaggregation for GRPO rollouts Brings up PD-disaggregated generation end-to-end on GB200: a replica's inference GPUs are split into context (prefill) and generation (decode) engines fronted by one OpenAI-compatible disagg server, which is the single URL NeMo Gym talks to. Engine plumbing. TrtllmGeneration plans engines per replica from trtllm_cfg.disaggregation (engine counts, per-role TP/EP overrides, routers, cache transceiver backend) and hands each worker its role. DisaggServerActor wraps TRT-LLM's OpenAIDisaggServer; trtllm_disagg_server.py adapts it to the NeMo Gym request shape. config.py gains the disaggregation schema, and build-custom-trtllm.sh plus the Dockerfile pick up UCX and NIXL so the cache transceiver is actually compiled in -- without them an engine aborts on the first KV transfer. Six bring-up fixes, each of which silently broke the path rather than failing loudly: - DisaggServerActor ran in the driver's environment and died with ModuleNotFoundError: tensorrt_llm. Give it the engine workers' interpreter, which RayWorkerGroup now exposes as py_executable. - OpenAIDisaggServer builds a prometheus MultiProcessCollector in register_routes(), which raises unless PROMETHEUS_MULTIPROC_DIR is set. TRT-LLM's own entrypoint calls set_prometheus_multiproc_dir() first; we construct the server directly, so call it too. - The middleware stripping NeMo Gym's vLLM-only request fields was a Starlette BaseHTTPMiddleware, which hands the downstream app its own captured receive channel, so reassigning request._receive never reached FastAPI's validation and every request 400'd on extra_forbidden. Rewritten as raw ASGI. - Aggregated serving lost its rollout fields: they moved off the message onto declared response fields for the disagg path, but with no disagg server to re-attach them NeMo Gym silently dropped every assistant turn. Attach them on the message when no disaggregation is in play. - Under disaggregation the unit that must stay inside one NVLink domain is the replica -- its context and generation engines exchange KV every turn -- not the engine. Sizing gpus_per_instance by the engine yielded nodes_per_instance=1 and skipped domain pinning entirely. - Per-node placement groups were consumed in creation order, so two adjacent pg_idx values could sit in different NVLink domains and split a replica across the fabric. Consume them in topology order instead. The engine HTTP server also moves from asyncio.to_thread(llm.generate) to llm.generate_async: the blocking API parks a worker thread per in-flight request, capping concurrency at the default executor size rather than at the engine's scheduler. Error fidelity. OpenAIDisaggServer._handle_exception only re-raises HTTPException, so a 4xx from an engine arrives as an aiohttp.ClientResponseError, falls into the catch-all, and reaches the caller as 500. The aggregated server returns it as a 4xx, and Gym accounts for the two classes differently -- which would make masking statistics incomparable between the aggregated and disaggregated paths, the exact comparison this work exists to support. Re-raise 400-499 with the original status; 5xx still goes to super(). Empty rollouts. A Gym rollout can return without a single assistant turn (the agent stalls before its first completion and Gym's wall-clock timeout kills it). That raised ValueError, taking the run down over one sample and losing the step's other 127 rollouts. NRL_SKIP_FAILING_EMPTY_ROLLOUT gates it: the default "0" keeps the raise, since the usual causes are misconfigurations worth surfacing; "1" stands the sample up as prompt-only and masks it out of the loss, reported as train/num_masked_seqs_by_empty_rollout. That count overlaps num_mask_sample_filtered by design and must not be summed with it -- num_valid_samples stays authoritative. Note there is no circuit breaker on a sustained rate. Profiling. Under disaggregation every engine runs the same worker class, so nsys reports differed only by %p pid and matching a trace to the context or generation side meant grepping the driver log. The -o filename now carries the role and an ordinal (_context0, _context1, _generation0), appended rather than prefixed so the report names documented in docs/nsys-profiling.md stay prefix-matchable. Also bumps TRT-LLM to 1.3.0rc24. Signed-off-by: shuyixiong <219646547+shuyixiong@users.noreply.github.com> (cherry picked from commit a8f1ad377066ccd4e091382c325aea8c342c7e23) --- docker/Dockerfile | 64 ++ examples/configs/grpo_math_1B_trtllm.yaml | 25 + nemo_rl/algorithms/grpo.py | 80 ++- nemo_rl/distributed/virtual_cluster.py | 58 ++ nemo_rl/distributed/worker_groups.py | 28 + nemo_rl/environments/nemo_gym.py | 103 ++- nemo_rl/experience/rollouts.py | 9 + nemo_rl/models/generation/trtllm/config.py | 49 ++ .../generation/trtllm/trtllm_disagg_server.py | 464 +++++++++++++ .../generation/trtllm/trtllm_generation.py | 617 +++++++++++++++--- .../generation/trtllm/trtllm_http_server.py | 165 ++++- .../generation/trtllm/trtllm_worker_async.py | 169 ++++- tools/build-custom-trtllm.sh | 34 +- 13 files changed, 1732 insertions(+), 133 deletions(-) create mode 100644 nemo_rl/models/generation/trtllm/trtllm_disagg_server.py 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/examples/configs/grpo_math_1B_trtllm.yaml b/examples/configs/grpo_math_1B_trtllm.yaml index d377cbda4ed..cf28abd4565 100644 --- a/examples/configs/grpo_math_1B_trtllm.yaml +++ b/examples/configs/grpo_math_1B_trtllm.yaml @@ -17,6 +17,31 @@ 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 + # Per-role overrides merged over trtllm_cfg (TP, the MoE split, and any + # other trtllm_cfg key). Uncomment to give the roles different shapes. + # ctx_trtllm_kwargs: + # tensor_parallel_size: 4 + # gen_trtllm_kwargs: + # tensor_parallel_size: 2 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..81a04929c66 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 @@ -2383,6 +2413,42 @@ def _apply_mask_sample_filter(repeated_batch: BatchedDataDict[DatumSpec]) -> int return num_masked +def _apply_empty_rollout_filter(repeated_batch: BatchedDataDict[DatumSpec]) -> int: + """Zero loss_multiplier where the rollout produced nothing, and count it. + + NemoGym stands a rollout that returned no assistant turn up as a + prompt-only sample so one dead rollout cannot fail the step. Such a sample + already contributes 0 to the loss -- it has no trainable tokens -- but + without zeroing loss_multiplier it still counts toward num_valid_samples, + which would then overstate how much of the batch actually trained. + + The returned count answers "how many rollouts came back empty", which is a + statement about generation health, and it deliberately counts every empty + rollout rather than only the ones this call was first to zero. The masking + metrics are attribution, not a partition: with + env.should_mask_flagged_samples on, Gym flags an agent that timed out + before its first completion, so the same sample is counted here and in + num_mask_sample_filtered. Zeroing is idempotent so the loss is unaffected, + but the counts overlap and must not be summed -- num_valid_samples + (sample_mask.sum()) is the one authoritative figure for how much of the + batch trained. + """ + if "empty_rollout" not in repeated_batch: + return 0 + + loss_multiplier = repeated_batch["loss_multiplier"].clone() + empty_rollout = repeated_batch["empty_rollout"] + + if isinstance(empty_rollout, list): + empty_rollout = torch.tensor(empty_rollout, dtype=torch.bool) + empty_rollout_bool = empty_rollout.bool() + + num_masked = int(empty_rollout_bool.sum().item()) + loss_multiplier[empty_rollout_bool] = 0 + repeated_batch["loss_multiplier"] = loss_multiplier + return num_masked + + def _should_log_nemo_gym_responses(master_config: MasterConfig) -> bool: """Whether NeMo Gym is responsible for full response logging. @@ -3397,6 +3463,10 @@ def grpo_train( num_mask_sample_filtered = _apply_mask_sample_filter(repeated_batch) metrics["num_mask_sample_filtered"] = num_mask_sample_filtered + metrics["num_masked_seqs_by_empty_rollout"] = ( + _apply_empty_rollout_filter(repeated_batch) + ) + add_grpo_token_loss_masks_and_generation_logprobs( repeated_batch["message_log"] ) @@ -5191,6 +5261,9 @@ def _flush_collector_telemetry() -> None: num_mask_sample_filtered = _apply_mask_sample_filter( repeated_batch ) + num_masked_seqs_by_empty_rollout = _apply_empty_rollout_filter( + repeated_batch + ) # Add loss mask to each message # Only unmask assistant messages that were actually generated (have generation_logprobs), @@ -5574,6 +5647,7 @@ def _flush_collector_telemetry() -> None: "loss": train_results["loss"].numpy(), "reward": rewards.numpy(), "num_mask_sample_filtered": num_mask_sample_filtered, + "num_masked_seqs_by_empty_rollout": num_masked_seqs_by_empty_rollout, "grad_norm": train_results["grad_norm"].numpy(), "mean_prompt_length": repeated_batch["length"].numpy(), "total_num_tokens": input_lengths.numpy(), diff --git a/nemo_rl/distributed/virtual_cluster.py b/nemo_rl/distributed/virtual_cluster.py index b3c4f96af94..7802aff3813 100644 --- a/nemo_rl/distributed/virtual_cluster.py +++ b/nemo_rl/distributed/virtual_cluster.py @@ -1187,6 +1187,64 @@ 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/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index 2c7c2235853..fffd21fd430 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -1098,30 +1098,99 @@ def _postprocess_nemo_gym_to_nemo_rl_result( output_item_dict["prompt_str"] = prompt_str output_item_dict["generation_str"] = generation_str - if not nemo_rl_message_log: + is_empty_rollout = not nemo_rl_message_log + if is_empty_rollout: + # A rollout can end without producing a single assistant turn: the + # agent stalls before its first completion (Gym's wall-clock timeout + # then kills it, `agent_timed_out`), or every output item was a + # reasoning/tool-call item. + # + # NRL_SKIP_FAILING_EMPTY_ROLLOUT decides what that costs. Default + # "0" fails the run, because the usual causes -- a first-turn prompt + # over max_model_len, or generation engines that are broken -- are + # misconfigurations worth surfacing loudly rather than silently + # training around. Set it to "1" on long runs, where losing a step's + # other 127 rollouts to one bad sample costs more than the sample is + # worth; the sample is then stood up as prompt-only and masked out + # of the loss. Note this trades a crash for a run that can keep + # going while producing nothing: watch + # train/num_masked_seqs_by_empty_rollout, since there is no circuit + # breaker on a sustained rate. input_messages = nemo_gym_result["responses_create_params"]["input"] + prompt_error: Optional[Exception] = None try: prompt_token_ids = tokenizer.apply_chat_template( input_messages, tokenize=True ) - prompt_len_str = f"{len(prompt_token_ids)} tokens" except Exception as e: - prompt_len_str = ( - f"" - ) + # An agent that died this early can leave `input` malformed, so + # the prompt is not always recoverable. + prompt_error = e + prompt_token_ids = None + output_item_types = [ o.get("type") for o in nemo_gym_result["response"]["output"] ] - raise ValueError( - f"NeMo Gym returned a result with no generation data. " - f"Possible causes: (1) the prompt for the first turn already exceeds the vLLM max_model_len, " - f"so vLLM rejected the request before any tokens could be generated; " - f"(2) all response output items were reasoning/tool-call items with no assistant generation.\n" - f" Prompt length: {prompt_len_str}.\n" - f" response.output item types ({len(output_item_types)} items): {output_item_types}.\n" - f" → If (1): increase `policy.max_total_sequence_length` and `policy.generation.vllm_cfg.max_model_len` " - f"above the prompt length above.\n" - f" → If (2): inspect why no assistant content was produced for this rollout." + + if os.environ.get("NRL_SKIP_FAILING_EMPTY_ROLLOUT", "0") != "1": + prompt_len_str = ( + f"{len(prompt_token_ids)} tokens" + if prompt_error is None + else f"" + ) + raise ValueError( + f"NeMo Gym returned a result with no generation data. " + f"Possible causes: (1) the prompt for the first turn already exceeds the vLLM max_model_len, " + f"so vLLM rejected the request before any tokens could be generated; " + f"(2) all response output items were reasoning/tool-call items with no assistant generation.\n" + f" Prompt length: {prompt_len_str}.\n" + f" response.output item types ({len(output_item_types)} items): {output_item_types}.\n" + f" → If (1): increase `policy.max_total_sequence_length` and `policy.generation.vllm_cfg.max_model_len` " + f"above the prompt length above.\n" + f" → If (2): inspect why no assistant content was produced for this rollout.\n" + f" → To mask samples like this out of the loss instead of failing, " + f"set NRL_SKIP_FAILING_EMPTY_ROLLOUT=1." + ) + + # Stand the sample up as prompt-only. With no assistant message it + # carries no trainable tokens, so token_loss_mask is empty and it + # contributes exactly 0 to the loss (masked_mean normalizes by the + # batch-global token count, so no division by zero). Its reward + # still reaches the group baseline, which is how every other masked + # sample behaves here -- overlong_filtering, Gym's mask_sample, and + # seq_logprob_error masking all zero the loss while leaving the + # reward in calculate_baseline_and_std_per_prompt (grpo.py passes + # torch.ones_like(rewards) as its valid_mask). + if prompt_error is not None: + # One pad token keeps the tensor typed and non-empty; nothing + # reads its value because the message is not trainable. + fallback_id = tokenizer.pad_token_id or tokenizer.eos_token_id or 0 + prompt_token_ids = [fallback_id] + print( + "NeMo Gym returned a rollout with no generation data and an " + f"unusable prompt ({type(prompt_error).__name__}: " + f"{prompt_error}); standing it up as a single-token " + "placeholder.", + file=sys.stderr, + ) + + print( + "NeMo Gym returned a result with no generation data; masking it " + "from the loss instead of failing the run " + "(NRL_SKIP_FAILING_EMPTY_ROLLOUT=1). response.output item " + f"types ({len(output_item_types)} items): {output_item_types}. " + "A run-wide rise in this message means rollouts are dying before " + "they generate -- check the generation engines rather than this " + "sample.", + file=sys.stderr, + ) + nemo_rl_message_log.append( + { + "role": "user", + "content": "", + "token_ids": torch.tensor(prompt_token_ids), + } ) if initial_multimodal_data_omitted: @@ -1139,6 +1208,10 @@ def _postprocess_nemo_gym_to_nemo_rl_result( "message_log": nemo_rl_message_log, "input_message_log": nemo_rl_message_log[:1], "full_result": nemo_gym_result, + # Surfaced as train/num_masked_seqs_by_empty_rollout: these samples + # stand in for a rollout that produced nothing, so a rising count + # means generation is failing, not that the policy is doing badly. + "empty_rollout": is_empty_rollout, } if not include_initial_multimodal_data: result["_initial_multimodal_data_omitted"] = initial_multimodal_data_omitted diff --git a/nemo_rl/experience/rollouts.py b/nemo_rl/experience/rollouts.py index a0688be3620..f07a3528bca 100644 --- a/nemo_rl/experience/rollouts.py +++ b/nemo_rl/experience/rollouts.py @@ -2978,6 +2978,15 @@ def _postprocess_single_nemo_gym_group( result["full_result"] for result in results ) + # Rollouts that returned no assistant turn at all. NemoGym stands these up + # as prompt-only samples so one dead rollout cannot fail the step. Carried + # unconditionally, unlike mask_sample: the cause is upstream of the policy + # -- a generation engine that stopped answering, not a bad trajectory -- so + # env.should_mask_flagged_samples has no say over it. + final_batch["empty_rollout"] = torch.tensor( + [bool(r.get("empty_rollout")) for r in results], dtype=torch.bool + ) + rollout_metrics.update(_effort_shaping_metrics(shaping)) rollout_metrics.update( diff --git a/nemo_rl/models/generation/trtllm/config.py b/nemo_rl/models/generation/trtllm/config.py index 10d3b16d205..21ad05dc565 100644 --- a/nemo_rl/models/generation/trtllm/config.py +++ b/nemo_rl/models/generation/trtllm/config.py @@ -17,6 +17,54 @@ from nemo_rl.models.generation.interfaces import GenerationConfig +class TrtllmDisaggArgs(TypedDict): + """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 + + # Engines per replica. The two are independent, so the P:D ratio is free. + num_context_engines: int + num_generation_engines: int + + # 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. + ctx_router: str # conversation | kv_cache_aware + gen_router: str # round_robin | load_balancing + + # Mapped onto TRT-LLM's CacheTransceiverConfig. + # DEFAULT | UCX | NIXL | MOONCAKE | MPI + cache_transceiver_backend: str + max_tokens_in_buffer: NotRequired[int] + + # 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. + ctx_trtllm_kwargs: NotRequired[dict[str, Any]] + gen_trtllm_kwargs: NotRequired[dict[str, Any]] + + class TrtllmSpecificArgs(TypedDict): tensor_parallel_size: int model_name: NotRequired[str] @@ -43,6 +91,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[TrtllmDisaggArgs] default_chat_template_kwargs: NotRequired[dict[str, Any]] # TRT-LLM's registered parser names: # "qwen3" -> Qwen3ToolParser (JSON format: {"name":..., "arguments":{...}}) 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..2e964e62238 --- /dev/null +++ b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py @@ -0,0 +1,464 @@ + +# 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 logging +import threading +from typing import Any, Optional + +import ray + +logger = logging.getLogger(__name__) + +__all__ = [ + "DisaggServerActor", + "DisaggServerActorImpl", + "build_config", + "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, +) -> 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 + ] + + return DisaggServerConfig( + server_configs=server_configs, + ctx_router_config=RouterConfig(type=ctx_router), + gen_router_config=RouterConfig(type=gen_router), + node_id=node_id, + ) + + +# Request fields NeMo-Gym sends for vLLM that TRT-LLM's ChatCompletionRequest +# does not declare. Its models are extra="forbid", so leaving them in means a +# 422 on every request. +_GYM_ONLY_REQUEST_FIELDS = ("return_tokens_as_token_ids", "return_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"}) + + +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) -> None: + self.app = app + + 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 + + import json + + 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 + 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) + body = json.dumps(payload).encode() + # Content-Length must follow the body; 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" + ] + 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: + 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) + + # -------------------------------------------------------------- # + # 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 + return self._attach_rollout_fields(response) + + return wrapper + + @staticmethod + 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(self, response: JSONResponse) -> JSONResponse: + """Re-attach the fields NeMo-Gym reads off ``choices[].message``.""" + import json + + payload = json.loads(response.body) + choices = payload.get("choices") or [] + if not choices: + return response + + choice = choices[0] + message = choice.get("message") + if not isinstance(message, dict): + return response + + if payload.get("prompt_token_ids") is not None: + message["prompt_token_ids"] = payload["prompt_token_ids"] + + token_ids = self._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 JSONResponse(content=payload, status_code=response.status_code) + + return OpenAIDisaggServerAdaptor + + +def start_server( + config: Any, + host: str = "0.0.0.0", + port: int = 0, + req_timeout_secs: int = 1800, +) -> "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, + ) + + def _run() -> None: + # OpenAIDisaggServer.__call__ is a coroutine that runs uvicorn, so the + # thread needs its own event loop. + 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) -> None: + self._replica_idx = replica_idx + self._thread = None + self._server = None + self._base_url: Optional[str] = None + + def start( + self, + ctx_addrs: list[tuple[str, int]], + gen_addrs: list[tuple[str, int]], + *, + ctx_router: str, + gen_router: str, + ) -> str: + """Serve this replica and return the URL NeMo-Gym will talk to.""" + if self._base_url is not None: + return self._base_url + + config = build_config( + ctx_addrs, + gen_addrs, + node_id=self._replica_idx, + ctx_router=ctx_router, + gen_router=gen_router, + ) + self._thread, self._base_url, self._server = start_server(config) + wait_ready(self._base_url) + + logger.info( + "disagg server for replica %d ready at %s", + 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 time + + 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..d1e5918eba8 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_generation.py +++ b/nemo_rl/models/generation/trtllm/trtllm_generation.py @@ -50,16 +50,47 @@ 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 = trtllm_cfg.get("disaggregation") or {} + engine_tp = trtllm_cfg["tensor_parallel_size"] + if disagg.get("enabled"): + engine_tp = max( + int((disagg.get(f"{role}_trtllm_kwargs") or {}).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.get("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 +110,57 @@ 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})." + + # 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.get("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 +168,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 +204,55 @@ 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.get("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 +260,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 +268,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 +283,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 +295,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: str) -> 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.get(f"{role}_trtllm_kwargs") or {} + return {**self.cfg["trtllm_cfg"], **overrides} + + def _plan_engines(self, world_size: int) -> tuple[list[str], list[int], int]: + """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.get("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 = int(disagg["num_context_engines"]) + num_gen = int(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.get("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 = "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.get(f"{prefix}_trtllm_kwargs") or {}), + } + + 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 +435,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 +472,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 +480,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 +533,176 @@ def _get_tied_worker_bundle_indices( ) return tied_groups + # ------------------------------------------------------------------ # + # PD disaggregation + # ------------------------------------------------------------------ # + + @property + def _disagg_cfg(self) -> dict[str, Any]: + return self.cfg["trtllm_cfg"].get("disaggregation") or {} + + 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.get("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 = int(disagg["num_context_engines"]) + num_gen = int(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." + ) + + addrs = self._report_engine_addrs() + missing = [i for i, a in enumerate(addrs) if not a] + assert not missing, f"engines {missing} reported no HTTP address" + + from ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy + + from nemo_rl.models.generation.trtllm.trtllm_disagg_server import ( + DisaggServerActor, + ) + + self._disagg_actors = [] + futures = [] + for replica_idx in range(self.num_replicas): + base = replica_idx * per_replica + # 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. Soft-pinned to the node holding its + # first context engine so routing hops stay local when they can. + actor = DisaggServerActor.options( + scheduling_strategy=NodeAffinitySchedulingStrategy( + node_id=addrs[base]["node_id"], soft=True + ), + name=f"trtllm_disagg_server_{replica_idx}", + # 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: No module named 'tensorrt_llm'. + # The venv already exists on every node -- the worker group + # created it before these actors are spawned. + runtime_env={"py_executable": self.worker_group.py_executable}, + ).remote(replica_idx) + self._disagg_actors.append(actor) + + futures.append( + actor.start.remote( + 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"], + ) + ) + + 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"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 OpenAI-compatible base URL 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.get("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 +721,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 +750,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 +800,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 +811,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 +822,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 +845,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 +862,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 +872,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 +882,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 +891,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 +905,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 +915,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 +962,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..2125c01e3f1 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_http_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_http_server.py @@ -13,8 +13,17 @@ # 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 logging @@ -33,6 +42,65 @@ logger = logging.getLogger(__name__) +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 @@ -157,6 +225,25 @@ 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. + 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"): @@ -218,6 +305,23 @@ async def chat_completions(request: Request): 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. + 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) + max_tokens_requested = ( body.get("max_tokens") or body.get("max_completion_tokens") or max_seq_len ) @@ -248,6 +352,7 @@ async def chat_completions(request: Request): output = await llm.generate_async( {"prompt_token_ids": adj_prompt}, sampling_params=sampling, + disaggregated_params=disagg_params, ) except RequestError as e: err = str(e) @@ -259,6 +364,16 @@ 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 +422,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 +429,44 @@ 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, + # Generated token ids. ChatCompletionResponseChoice needs + # the matching field upstream (CompletionResponseChoice + # already has it) or the disagg server rejects this. + "token_ids": gen_token_ids, + } ], + # 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), @@ -338,10 +475,18 @@ async def chat_completions(request: Request): } 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. + as_ids = bool(body.get("return_tokens_as_token_ids")) 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..91c6f9f08df 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,61 @@ 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 +190,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 +226,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 +245,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 +255,42 @@ 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" + ] + 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 +323,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 +361,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 +401,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 +435,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/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 \ From 321acc96c07a5cd5ae5a968d3369f259b6e96529 Mon Sep 17 00:00:00 2001 From: Zongfei Jing <20381269+zongfeijing@users.noreply.github.com> Date: Tue, 1 Sep 2026 16:25:30 -0700 Subject: [PATCH 02/12] disagg: plumb transceiver runtime, bounce size and transfer timeout Three CacheTransceiverConfig knobs the recipe could not reach, each forwarded only when set so TRT-LLM's own defaults hold otherwise: - cache_transceiver_runtime: "auto" silently 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 for its recurrent-state handoff -- so the recipe must be able to force it. - kv_cache_bounce_size_mb: coalesces a request's scattered per-block KV into one contiguous fabric-VMM buffer and a single multi-rail NIXL write, sidestepping per-block registration failures. - kv_transfer_timeout_ms: TRT-LLM's 60 s default is tuned for short prompts at low concurrency; at high rollout concurrency bulk ctx-side timeouts feed a cancel/retry churn that stresses the transceiver, so large multi-turn workloads want a much larger value. Co-Authored-By: Claude Fable 5 (cherry picked from commit 407f3832a5a87559e10fc701c901da57e7a53d63) --- nemo_rl/models/generation/trtllm/config.py | 23 +++++++++++++++++++ .../generation/trtllm/trtllm_worker_async.py | 17 ++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/nemo_rl/models/generation/trtllm/config.py b/nemo_rl/models/generation/trtllm/config.py index 21ad05dc565..0138e5a476e 100644 --- a/nemo_rl/models/generation/trtllm/config.py +++ b/nemo_rl/models/generation/trtllm/config.py @@ -54,8 +54,31 @@ class TrtllmDisaggArgs(TypedDict): # Mapped onto TRT-LLM's CacheTransceiverConfig. # DEFAULT | UCX | NIXL | MOONCAKE | MPI cache_transceiver_backend: str + + # "CPP" | "PYTHON" | "auto". 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: NotRequired[str] + + # 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: NotRequired[int] max_tokens_in_buffer: NotRequired[int] + # 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: NotRequired[int] + # 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 diff --git a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py index 91c6f9f08df..b6dacb50b79 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py +++ b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py @@ -285,6 +285,23 @@ def __init__( 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 ) From d6d1e545a27cad99ea687f340b5cdb8bd646c9d4 Mon Sep 17 00:00:00 2001 From: Zongfei Jing <20381269+zongfeijing@users.noreply.github.com> Date: Tue, 1 Sep 2026 16:25:56 -0700 Subject: [PATCH 03/12] disagg: shard per-replica frontends and instrument the full request path At high rollout concurrency the single DisaggServerActor (one uvicorn process) saturates around 17 turns/s and becomes the replica's ceiling: conversations pile up ahead of the engines while GPUs idle. This makes the frontend horizontally scalable and, since the same investigation needed to see where pre-generation time actually goes, adds permanent end-to-end timing stamps. Frontend sharding (num_frontend_workers, default 1 = old behavior): - trtllm_generation: replica x frontend actor fan-out; NeMo-Gym receives every frontend URL and its per-session client selection pins each conversation to one frontend (sticky routing is per-process state, so stickiness only holds within a frontend -- pinning sessions is what makes N processes correct). - Deterministic ports (frontend_base_port + idx on the pinned node) and serve-from-constructor: a Ray-restarted actor re-binds the same port, so the URL Gym holds stays valid across crashes. - Snowflake node_id = replica_idx * num_frontends + frontend_idx keeps ctx_request_id unique across frontends (process_id is hardwired 0 in-process and time.monotonic shares an origin per node). Per-request work moved off the hot single processes: - gen_tokids_ctxbytes / gen_strip_message_history: the gen leg carries b64 int32 token ids instead of a 30k-int JSON array plus the full message history it never reads. - frontend_tokenize: frontends render the chat template and tokenize via build_spliced_prompt_ids -- the exact pipeline the adapters use, now hoisted to module level as the single source of truth -- and attach prompt_token_ids_b64 to the ctx leg. Inbound payloads are normalized through ChatCompletionRequest.model_dump(exclude_unset) first: the raw Gym payload orders tool-JSON keys differently than the pydantic-normalized form the adapters see, and the chat template is sensitive to that order (turn-1 only; the splice covers later turns). Guarded by ctx-side shadow validation (NRL_TRTLLM_TOKENIZE_SHADOW_RATE) which re-tokenizes a sample on the adapter and logs any divergence with a decoded token window. - The supplied-ids short-circuit in the adapter route is hoisted before template rendering, so the gen leg no longer renders a template whose output it discards. Full-path timeline stamps (all gated NRL_TRTLLM_EMIT_TIMELINE_FIELDS): - Frontend ASGI middleware stamps receive/forward times as headers; the response wrapper surfaces them as nemo_fe_recv/fwd_ts_us. - The ctx adapter emits nemo_ctx_{arrival,queued,first_scheduled, done}_ts_us on the context leg; the disagg service relays them onto the final response (tekit carries the relay change), decomposing the previously-opaque pre-generation leg into frontend / ctx submit / ctx queue / prefill / KV-handoff segments per model call. Validated end to end at conc-512 (2048 rollouts, 30-turn agentic workload, 61k model calls): stamps present on 100% of calls, sub-leg sum identical to the parent leg, tokenize shadow divergence zero. Sharding the frontend (N=8) plus frontend tokenize took the equal-GPU disagg configuration from well behind the aggregated baseline to ahead of it. Co-Authored-By: Claude Fable 5 (cherry picked from commit 769aab4177d35a3601771fe21045b489fa8643d6) --- nemo_rl/models/generation/trtllm/config.py | 24 ++ .../generation/trtllm/trtllm_disagg_server.py | 237 ++++++++++++++++-- .../generation/trtllm/trtllm_generation.py | 98 +++++--- .../generation/trtllm/trtllm_http_server.py | 193 ++++++++++---- 4 files changed, 446 insertions(+), 106 deletions(-) diff --git a/nemo_rl/models/generation/trtllm/config.py b/nemo_rl/models/generation/trtllm/config.py index 0138e5a476e..809a07339b1 100644 --- a/nemo_rl/models/generation/trtllm/config.py +++ b/nemo_rl/models/generation/trtllm/config.py @@ -51,6 +51,30 @@ class TrtllmDisaggArgs(TypedDict): ctx_router: str # conversation | kv_cache_aware gen_router: str # round_robin | 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: NotRequired[int] + # 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: NotRequired[bool] + # 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: NotRequired[bool] + # 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: NotRequired[bool] + # Base port for the frontend workers' deterministic ports + # (base + frontend_idx on the pinned node). 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: NotRequired[int] + # Mapped onto TRT-LLM's CacheTransceiverConfig. # DEFAULT | UCX | NIXL | MOONCAKE | MPI cache_transceiver_backend: str diff --git a/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py index 2e964e62238..6575243698d 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py @@ -24,6 +24,7 @@ import logging import threading +import time from typing import Any, Optional import ray @@ -46,6 +47,8 @@ def build_config( 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. @@ -75,12 +78,18 @@ def build_config( for host, port in gen_addrs ] - return DisaggServerConfig( + 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 NeMo-Gym sends for vLLM that TRT-LLM's ChatCompletionRequest @@ -105,14 +114,19 @@ class passes the downstream app its own captured receive channel, so actually reads. """ - def __init__(self, app: Any) -> None: + 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 + import json chunks: list[bytes] = [] @@ -129,21 +143,49 @@ async def __call__(self, scope: Any, receive: Any, send: Any) -> None: 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; 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" - ] + # 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} + scope = {**scope, "headers": headers} delivered = False @@ -188,6 +230,7 @@ class OpenAIDisaggServerAdaptor(OpenAIDisaggServer): """ 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() @@ -221,7 +264,9 @@ def _install_request_adaptor(self) -> None: # 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) + self.app.add_middleware( + _DropGymOnlyRequestFields, tokenize_fn=self._tokenize_fn + ) # -------------------------------------------------------------- # # Errors @@ -263,7 +308,7 @@ async def wrapper(req: request_type, raw_req: Request) -> Response: # type: ign response = await inner(req, raw_req) if req.stream or not isinstance(response, JSONResponse): return response - return self._attach_rollout_fields(response) + return self._attach_rollout_fields(response, raw_req) return wrapper @@ -288,19 +333,26 @@ def _generation_token_ids(choice: dict[str, Any]) -> Optional[list[int]]: ids.append(int(token[len(_TOKEN_ID_PREFIX) :])) return ids or None - def _attach_rollout_fields(self, response: JSONResponse) -> JSONResponse: + def _attach_rollout_fields( + self, response: JSONResponse, raw_req: "Request | None" = None + ) -> JSONResponse: """Re-attach the fields NeMo-Gym reads off ``choices[].message``.""" import json payload = json.loads(response.body) + choices = payload.get("choices") or [] if not choices: - return response + return JSONResponse( + content=payload, status_code=response.status_code + ) choice = choices[0] message = choice.get("message") if not isinstance(message, dict): - return response + return JSONResponse( + content=payload, status_code=response.status_code + ) if payload.get("prompt_token_ids") is not None: message["prompt_token_ids"] = payload["prompt_token_ids"] @@ -320,11 +372,65 @@ def _attach_rollout_fields(self, response: JSONResponse) -> JSONResponse: 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 transformers import AutoConfig, AutoTokenizer + + from nemo_rl.models.generation.trtllm.trtllm_http_server import ( + build_spliced_prompt_ids, + ) + + from tensorrt_llm.serve.openai_protocol import ChatCompletionRequest + + 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 @@ -357,11 +463,28 @@ def start_server( 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. + # 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) @@ -385,36 +508,96 @@ class DisaggServerActorImpl: so it can be exercised without Ray. """ - def __init__(self, replica_idx: int) -> None: + 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 - - def start( + # 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, - ) -> str: - """Serve this replica and return the URL NeMo-Gym will talk to.""" - if self._base_url is not None: - return self._base_url - + 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, + 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, ) - self._thread, self._base_url, self._server = start_server(config) + 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, + ctx_addrs: Optional[list[tuple[str, int]]] = None, + gen_addrs: Optional[list[tuple[str, int]]] = None, + *, + ctx_router: str = "conversation", + gen_router: str = "load_balancing", + gen_tokids_ctxbytes: bool = False, + gen_strip_message_history: bool = False, + ) -> str: + """Ensure the server is up and return the URL NeMo-Gym will talk to.""" + if self._base_url is None: + assert ctx_addrs is not None and gen_addrs is not None, ( + "start() needs the address pools unless they were passed to " + "the constructor via serve_args" + ) + self._start_serving( + ctx_addrs, + gen_addrs, + ctx_router=ctx_router, + gen_router=gen_router, + gen_tokids_ctxbytes=gen_tokids_ctxbytes, + gen_strip_message_history=gen_strip_message_history, + ) wait_ready(self._base_url) logger.info( - "disagg server for replica %d ready at %s", + "disagg frontend %d/%d for replica %d ready at %s", + self._frontend_idx, + self._num_frontends, self._replica_idx, self._base_url, ) diff --git a/nemo_rl/models/generation/trtllm/trtllm_generation.py b/nemo_rl/models/generation/trtllm/trtllm_generation.py index d1e5918eba8..9a13985e3a0 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_generation.py +++ b/nemo_rl/models/generation/trtllm/trtllm_generation.py @@ -593,48 +593,80 @@ def _start_disagg_servers(self) -> list[Optional[str]]: DisaggServerActor, ) + n_fe = int(disagg.get("num_frontend_workers") or 1) + 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); frontends on one node get base+frontend_idx. + base_port = int(disagg.get("frontend_base_port") or 17300) + self._disagg_actors = [] futures = [] for replica_idx in range(self.num_replicas): base = replica_idx * per_replica - # 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. Soft-pinned to the node holding its - # first context engine so routing hops stay local when they can. - actor = DisaggServerActor.options( - scheduling_strategy=NodeAffinitySchedulingStrategy( - node_id=addrs[base]["node_id"], soft=True + 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=bool(disagg.get("gen_tokids_ctxbytes", False)), + gen_strip_message_history=bool( + disagg.get("gen_strip_message_history", False) + ), + frontend_tokenize=bool(disagg.get("frontend_tokenize", False)), + model_name=self.cfg["model_name"], + default_chat_template_kwargs=self.cfg["trtllm_cfg"].get( + "default_chat_template_kwargs" ), - name=f"trtllm_disagg_server_{replica_idx}", - # 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: No module named 'tensorrt_llm'. - # The venv already exists on every node -- the worker group - # created it before these actors are spawned. - runtime_env={"py_executable": self.worker_group.py_executable}, - ).remote(replica_idx) - self._disagg_actors.append(actor) - - futures.append( - actor.start.remote( - 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"], - ) ) + 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 + 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"({num_ctx} context + {num_gen} generation) engines, " + f"{n_fe} frontend worker(s) each; " f"disagg servers: {self._disagg_server_urls}", flush=True, ) @@ -683,7 +715,7 @@ def _report_engine_addrs(self) -> list[Optional[dict[str, Any]]]: return ray.get(futures) def _report_dp_openai_server_base_urls(self) -> list[Optional[str]]: - """One OpenAI-compatible base URL per DP shard, in shard order. + """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 diff --git a/nemo_rl/models/generation/trtllm/trtllm_http_server.py b/nemo_rl/models/generation/trtllm/trtllm_http_server.py index 2125c01e3f1..1259008f986 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_http_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_http_server.py @@ -26,7 +26,10 @@ rebuilding them, so the sequence matches the KV that was transferred. """ +import asyncio import logging +import os +import random import threading import time import uuid @@ -42,6 +45,78 @@ logger = logging.getLogger(__name__) +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], @@ -180,8 +255,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( @@ -263,52 +336,15 @@ 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. + # 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. @@ -319,8 +355,65 @@ async def chat_completions(request: Request): 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 @@ -480,7 +573,15 @@ async def chat_completions(request: Request): # 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. - as_ids = bool(body.get("return_tokens_as_token_ids")) + # 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": [ { From 8d11effd85beae0ad458d55dc3377bd24597e5b3 Mon Sep 17 00:00:00 2001 From: Zongfei Jing <20381269+zongfeijing@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:30:50 -0700 Subject: [PATCH 04/12] disagg: relay the conversation identity to the engine and expose ctx compute diagnostics The conversation-affine ADP router (kv_cache_routing_conversation_affinity) pins a conversation's turns to the rank that holds its prefix, but it only sees the id when the request carries ConversationParams. The HTTP adapter never forwarded one, so under attention-DP every turn landed on an effectively random rank: measured at 2P-DEP8 / conc 512 only 17-26% of turns found the previous turn's KV on the serving rank (about 1/DEP) and the ctx engine re-prefilled most of the history every turn. Read the id from the body's conversation_params (the Gym model proxy sends its session id there; one session per rollout) or, failing that, from the id the disagg service stamps onto disaggregated_params for its ctx/gen legs, and pass ConversationParams(conversation_id) to llm.generate_async. No id means no affinity, i.e. the previous behaviour. With the id: 96-97% of turns on the rank holding their prefix, ctx prefill tokens -66% to -88%. Also surface the context leg's compute accounting on the ctx response and relay it with the other ctx stamps (nemo_ctx_computed_tokens, nemo_ctx_first_begin, nemo_ctx_num_chunks): the engine's cached_tokens overshoots the real reuse point by a few blocks, so computed tokens and the first chunk's begin position are what the reuse analysis needs. These live on the RequestOutput (the per-choice CompletionOutput only forwards cached_tokens), hence the extra request_output argument. Co-Authored-By: Claude Fable 5.1 (cherry picked from commit 14690fb5fa4aa246006fbe663584ddc643a1e388) --- .../models/generation/trtllm/trtllm_http_server.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/nemo_rl/models/generation/trtllm/trtllm_http_server.py b/nemo_rl/models/generation/trtllm/trtllm_http_server.py index 1259008f986..744c23fb38d 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_http_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_http_server.py @@ -303,6 +303,11 @@ async def chat_completions(request: Request): # -- 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. + _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 ( @@ -442,10 +447,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) From 62716384df153e74dc163cbd89a02b1a44ce2595 Mon Sep 17 00:00:00 2001 From: Zongfei Jing <20381269+zongfeijing@users.noreply.github.com> Date: Mon, 7 Sep 2026 16:43:55 -0700 Subject: [PATCH 05/12] ray.sub: opt-in NRL_CONTAINER_ENV_FORWARD to force host env over image ENV Comma-separated variable names get --container-env on every srun so the host value beats the image's baked ENV (pyxis lets the image win by default). First use: UCX_NET_DEVICES=all for NIXL KV transfer over the RDMA HCAs (the image bakes the management NIC os_p1s0) when running PD-disaggregated generation under this launcher. (cherry picked from commit f68ebd8a71e630efe576ba66f5f2835b1795fc4b) --- ray.sub | 8 ++++++++ 1 file changed, 8 insertions(+) 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 From 29ab8b87832b431820a17847af4f59d314dd82e0 Mon Sep 17 00:00:00 2001 From: Zongfei Jing <20381269+zongfeijing@users.noreply.github.com> Date: Tue, 8 Sep 2026 07:24:41 -0700 Subject: [PATCH 06/12] disagg: strip the vLLM-only required_prefix_token_ids field in the frontend middleware The MLPerf warmup sends required_prefix_token_ids (the vLLM server's prefix override extension) with every synthetic request. TRT-LLM's OpenAIDisaggServer validates bodies against extra="forbid" models, so the field turned each warmup request into a 400 and the warmup exception then took the driver down before run_start. The TRT-LLM engine adapter derives the on-policy prefix from the assistant messages and ignores the top-level field, so dropping it in the disagg frontend (like the Dynamo wrapper does) keeps both paths identical. (cherry picked from commit 46cd1a1f744cb4ae2fba00617ed3c5a375ed9bb0) --- .../generation/trtllm/trtllm_disagg_server.py | 13 +- .../trtllm/test_trtllm_disagg_server.py | 135 ++++++++++++++++++ 2 files changed, 145 insertions(+), 3 deletions(-) create mode 100644 tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py diff --git a/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py index 6575243698d..25466ff21c0 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py @@ -92,10 +92,17 @@ def build_config( return config -# Request fields NeMo-Gym sends for vLLM that TRT-LLM's ChatCompletionRequest +# 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 on every request. -_GYM_ONLY_REQUEST_FIELDS = ("return_tokens_as_token_ids", "return_token_ids") +# 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:" 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..910d6d2c22d --- /dev/null +++ b/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py @@ -0,0 +1,135 @@ +# 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 + +from nemo_rl.models.generation.trtllm.trtllm_disagg_server import ( + _GYM_ONLY_REQUEST_FIELDS, + _DropGymOnlyRequestFields, +) + + +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 From cc12c1497a07db5fba37c8a7d2ed8242fd3f5ca3 Mon Sep 17 00:00:00 2001 From: shuyixiong <219646547+shuyixiong@users.noreply.github.com> Date: Thu, 10 Sep 2026 17:30:23 -0700 Subject: [PATCH 07/12] revert(trtllm): drop the empty-rollout masking from the disagg base The disagg base commit bundled a second, independent change: NRL_SKIP_FAILING_EMPTY_ROLLOUT, which turns "NeMo Gym returned a rollout with no assistant turn" from a hard failure into a prompt-only sample masked out of the loss. That is an error-handling policy change, not part of prefill/decode disaggregation, and it changes behaviour for every backend. Take it out and restore the original raise. Removed: - nemo_rl/environments/nemo_gym.py -- restored verbatim; its entire delta over the branch point was this feature. An empty rollout raises ValueError again, with the original message (no NRL_SKIP_FAILING_EMPTY_ROLLOUT hint). - nemo_rl/experience/rollouts.py -- restored verbatim; likewise, the only delta was carrying the `empty_rollout` flag onto the batch. - nemo_rl/algorithms/grpo.py -- _apply_empty_rollout_filter and its two call sites (grpo_train, async_grpo_train) plus the train/num_masked_seqs_by_empty_rollout metric. Kept in grpo.py: the disagg-aware gpus_per_instance sizing, which is the part of that commit's grpo.py delta that does belong to disaggregation -- under PD the NVLink-domain unit is the replica, not the engine, and sizing by the engine skipped domain pinning entirely. Verified: no empty_rollout / NRL_SKIP_FAILING_EMPTY_ROLLOUT / num_masked_seqs_by_empty references remain, and the three modules compile. Co-Authored-By: Claude Opus 5 --- nemo_rl/algorithms/grpo.py | 44 ------------- nemo_rl/environments/nemo_gym.py | 103 +++++-------------------------- nemo_rl/experience/rollouts.py | 9 --- 3 files changed, 15 insertions(+), 141 deletions(-) diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index 81a04929c66..bbc35ac2496 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -2413,42 +2413,6 @@ def _apply_mask_sample_filter(repeated_batch: BatchedDataDict[DatumSpec]) -> int return num_masked -def _apply_empty_rollout_filter(repeated_batch: BatchedDataDict[DatumSpec]) -> int: - """Zero loss_multiplier where the rollout produced nothing, and count it. - - NemoGym stands a rollout that returned no assistant turn up as a - prompt-only sample so one dead rollout cannot fail the step. Such a sample - already contributes 0 to the loss -- it has no trainable tokens -- but - without zeroing loss_multiplier it still counts toward num_valid_samples, - which would then overstate how much of the batch actually trained. - - The returned count answers "how many rollouts came back empty", which is a - statement about generation health, and it deliberately counts every empty - rollout rather than only the ones this call was first to zero. The masking - metrics are attribution, not a partition: with - env.should_mask_flagged_samples on, Gym flags an agent that timed out - before its first completion, so the same sample is counted here and in - num_mask_sample_filtered. Zeroing is idempotent so the loss is unaffected, - but the counts overlap and must not be summed -- num_valid_samples - (sample_mask.sum()) is the one authoritative figure for how much of the - batch trained. - """ - if "empty_rollout" not in repeated_batch: - return 0 - - loss_multiplier = repeated_batch["loss_multiplier"].clone() - empty_rollout = repeated_batch["empty_rollout"] - - if isinstance(empty_rollout, list): - empty_rollout = torch.tensor(empty_rollout, dtype=torch.bool) - empty_rollout_bool = empty_rollout.bool() - - num_masked = int(empty_rollout_bool.sum().item()) - loss_multiplier[empty_rollout_bool] = 0 - repeated_batch["loss_multiplier"] = loss_multiplier - return num_masked - - def _should_log_nemo_gym_responses(master_config: MasterConfig) -> bool: """Whether NeMo Gym is responsible for full response logging. @@ -3463,10 +3427,6 @@ def grpo_train( num_mask_sample_filtered = _apply_mask_sample_filter(repeated_batch) metrics["num_mask_sample_filtered"] = num_mask_sample_filtered - metrics["num_masked_seqs_by_empty_rollout"] = ( - _apply_empty_rollout_filter(repeated_batch) - ) - add_grpo_token_loss_masks_and_generation_logprobs( repeated_batch["message_log"] ) @@ -5261,9 +5221,6 @@ def _flush_collector_telemetry() -> None: num_mask_sample_filtered = _apply_mask_sample_filter( repeated_batch ) - num_masked_seqs_by_empty_rollout = _apply_empty_rollout_filter( - repeated_batch - ) # Add loss mask to each message # Only unmask assistant messages that were actually generated (have generation_logprobs), @@ -5647,7 +5604,6 @@ def _flush_collector_telemetry() -> None: "loss": train_results["loss"].numpy(), "reward": rewards.numpy(), "num_mask_sample_filtered": num_mask_sample_filtered, - "num_masked_seqs_by_empty_rollout": num_masked_seqs_by_empty_rollout, "grad_norm": train_results["grad_norm"].numpy(), "mean_prompt_length": repeated_batch["length"].numpy(), "total_num_tokens": input_lengths.numpy(), diff --git a/nemo_rl/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index fffd21fd430..2c7c2235853 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -1098,99 +1098,30 @@ def _postprocess_nemo_gym_to_nemo_rl_result( output_item_dict["prompt_str"] = prompt_str output_item_dict["generation_str"] = generation_str - is_empty_rollout = not nemo_rl_message_log - if is_empty_rollout: - # A rollout can end without producing a single assistant turn: the - # agent stalls before its first completion (Gym's wall-clock timeout - # then kills it, `agent_timed_out`), or every output item was a - # reasoning/tool-call item. - # - # NRL_SKIP_FAILING_EMPTY_ROLLOUT decides what that costs. Default - # "0" fails the run, because the usual causes -- a first-turn prompt - # over max_model_len, or generation engines that are broken -- are - # misconfigurations worth surfacing loudly rather than silently - # training around. Set it to "1" on long runs, where losing a step's - # other 127 rollouts to one bad sample costs more than the sample is - # worth; the sample is then stood up as prompt-only and masked out - # of the loss. Note this trades a crash for a run that can keep - # going while producing nothing: watch - # train/num_masked_seqs_by_empty_rollout, since there is no circuit - # breaker on a sustained rate. + if not nemo_rl_message_log: input_messages = nemo_gym_result["responses_create_params"]["input"] - prompt_error: Optional[Exception] = None try: prompt_token_ids = tokenizer.apply_chat_template( input_messages, tokenize=True ) + prompt_len_str = f"{len(prompt_token_ids)} tokens" except Exception as e: - # An agent that died this early can leave `input` malformed, so - # the prompt is not always recoverable. - prompt_error = e - prompt_token_ids = None - + prompt_len_str = ( + f"" + ) output_item_types = [ o.get("type") for o in nemo_gym_result["response"]["output"] ] - - if os.environ.get("NRL_SKIP_FAILING_EMPTY_ROLLOUT", "0") != "1": - prompt_len_str = ( - f"{len(prompt_token_ids)} tokens" - if prompt_error is None - else f"" - ) - raise ValueError( - f"NeMo Gym returned a result with no generation data. " - f"Possible causes: (1) the prompt for the first turn already exceeds the vLLM max_model_len, " - f"so vLLM rejected the request before any tokens could be generated; " - f"(2) all response output items were reasoning/tool-call items with no assistant generation.\n" - f" Prompt length: {prompt_len_str}.\n" - f" response.output item types ({len(output_item_types)} items): {output_item_types}.\n" - f" → If (1): increase `policy.max_total_sequence_length` and `policy.generation.vllm_cfg.max_model_len` " - f"above the prompt length above.\n" - f" → If (2): inspect why no assistant content was produced for this rollout.\n" - f" → To mask samples like this out of the loss instead of failing, " - f"set NRL_SKIP_FAILING_EMPTY_ROLLOUT=1." - ) - - # Stand the sample up as prompt-only. With no assistant message it - # carries no trainable tokens, so token_loss_mask is empty and it - # contributes exactly 0 to the loss (masked_mean normalizes by the - # batch-global token count, so no division by zero). Its reward - # still reaches the group baseline, which is how every other masked - # sample behaves here -- overlong_filtering, Gym's mask_sample, and - # seq_logprob_error masking all zero the loss while leaving the - # reward in calculate_baseline_and_std_per_prompt (grpo.py passes - # torch.ones_like(rewards) as its valid_mask). - if prompt_error is not None: - # One pad token keeps the tensor typed and non-empty; nothing - # reads its value because the message is not trainable. - fallback_id = tokenizer.pad_token_id or tokenizer.eos_token_id or 0 - prompt_token_ids = [fallback_id] - print( - "NeMo Gym returned a rollout with no generation data and an " - f"unusable prompt ({type(prompt_error).__name__}: " - f"{prompt_error}); standing it up as a single-token " - "placeholder.", - file=sys.stderr, - ) - - print( - "NeMo Gym returned a result with no generation data; masking it " - "from the loss instead of failing the run " - "(NRL_SKIP_FAILING_EMPTY_ROLLOUT=1). response.output item " - f"types ({len(output_item_types)} items): {output_item_types}. " - "A run-wide rise in this message means rollouts are dying before " - "they generate -- check the generation engines rather than this " - "sample.", - file=sys.stderr, - ) - nemo_rl_message_log.append( - { - "role": "user", - "content": "", - "token_ids": torch.tensor(prompt_token_ids), - } + raise ValueError( + f"NeMo Gym returned a result with no generation data. " + f"Possible causes: (1) the prompt for the first turn already exceeds the vLLM max_model_len, " + f"so vLLM rejected the request before any tokens could be generated; " + f"(2) all response output items were reasoning/tool-call items with no assistant generation.\n" + f" Prompt length: {prompt_len_str}.\n" + f" response.output item types ({len(output_item_types)} items): {output_item_types}.\n" + f" → If (1): increase `policy.max_total_sequence_length` and `policy.generation.vllm_cfg.max_model_len` " + f"above the prompt length above.\n" + f" → If (2): inspect why no assistant content was produced for this rollout." ) if initial_multimodal_data_omitted: @@ -1208,10 +1139,6 @@ def _postprocess_nemo_gym_to_nemo_rl_result( "message_log": nemo_rl_message_log, "input_message_log": nemo_rl_message_log[:1], "full_result": nemo_gym_result, - # Surfaced as train/num_masked_seqs_by_empty_rollout: these samples - # stand in for a rollout that produced nothing, so a rising count - # means generation is failing, not that the policy is doing badly. - "empty_rollout": is_empty_rollout, } if not include_initial_multimodal_data: result["_initial_multimodal_data_omitted"] = initial_multimodal_data_omitted diff --git a/nemo_rl/experience/rollouts.py b/nemo_rl/experience/rollouts.py index f07a3528bca..a0688be3620 100644 --- a/nemo_rl/experience/rollouts.py +++ b/nemo_rl/experience/rollouts.py @@ -2978,15 +2978,6 @@ def _postprocess_single_nemo_gym_group( result["full_result"] for result in results ) - # Rollouts that returned no assistant turn at all. NemoGym stands these up - # as prompt-only samples so one dead rollout cannot fail the step. Carried - # unconditionally, unlike mask_sample: the cause is upstream of the policy - # -- a generation engine that stopped answering, not a bad trajectory -- so - # env.should_mask_flagged_samples has no say over it. - final_batch["empty_rollout"] = torch.tensor( - [bool(r.get("empty_rollout")) for r in results], dtype=torch.bool - ) - rollout_metrics.update(_effort_shaping_metrics(shaping)) rollout_metrics.update( From cc71fe40f2b2f3f30225ad6c0081619de6289c98 Mon Sep 17 00:00:00 2001 From: shuyixiong <219646547+shuyixiong@users.noreply.github.com> Date: Mon, 21 Sep 2026 23:50:09 -0700 Subject: [PATCH 08/12] fix(trtllm): replica-index the disagg frontend port and harden its contract Review fixes for the PD disaggregation PR, scoped to nemo_rl/ and tests/. The build script, Dockerfile, docs, exemplar YAML and pyrefly.toml changes are left for a separate commit. The one behaviour change is the frontend port. `base_port + fe_idx` carried no replica term while the node pin (`addrs[base + fe_idx % per_replica]`) did, so with num_frontend_workers=1 every replica sharing a node asked for 17300. Nothing surfaced the collision: start_server computes base_url before binding and returns it unconditionally, uvicorn's bind failure calls sys.exit(1) which only unwinds the daemon thread, and wait_ready's /health GET is then answered 200 by the winner bound on 0.0.0.0. The losers advertised the winner's URL and their engines idled with no exception and no timeout -- at the exemplar's defaults on one 8-GPU node that is 4 replicas on one port and 6 of 8 GPUs idle for the whole run. Now `base_port + replica_idx * n_fe + fe_idx`, still deterministic per (replica, frontend) so a restarted actor re-binds the same port, plus a thread-liveness assert after wait_ready so a future bind failure cannot pass silently again. TrtllmDisaggArgs becomes TrtllmDisaggConfig(BaseModel, extra="allow"), per the config conventions: all 14 fields carry their default on the schema rather than at call sites (num_frontend_workers, frontend_base_port, gen_tokids_ctxbytes, gen_strip_message_history and frontend_tokenize each had one). The routers and cache_transceiver_backend become Literal -- ctx_router was typed str, so "round_robin" parsed fine and silently discarded the prefix affinity the context engines depend on. resolve_trtllm_disagg_config mirrors resolve_vllm_video_config. DisaggServerActorImpl.start() drops its parameters. Serving is the constructor's job -- Ray replays __init__ on actor restart but never method calls -- so the `if self._base_url is None` branch was unreachable, and also less capable than the real path: no frontend_tokenize, no model_name, and routers hardcoded over the configured values. Asserted instead of kept as a fallback that could only come up misconfigured. generation_token_ids and attach_rollout_fields move to module level so the disagg-vs-aggregated rollout-field contract is testable without importing tensorrt_llm; the adaptor keeps only the JSONResponse re-wrap. A miss there drops a turn's training data silently (nemo_gym.py skips a message with no generation_token_ids), so it should not be the one half that no test can reach. trtllm_http_server stops emitting choices[].token_ids unconditionally: on a build whose ChatCompletionResponseChoice does not declare the field, the write 400s at the disagg server's extra="forbid" re-validation -- which is exactly where the token_id:N logprobs fallback is meant to take over. Gated on a cached capability probe so the read and the write agree. Tests: test_trtllm_disagg_server.py had no pytest.mark.trtllm, so no L0 shard collected it and deleting required_prefix_token_ids from the strip list would have left CI green; test_trtllm_http_server.py had the same pre-existing gap. Adds middleware coverage for the tokenizer-failure fallback, the disaggregated_params guard, chunked body reassembly and non-JSON passthrough, plus rollout-field, disagg-config and _plan_engines cases. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: shuyixiong <219646547+shuyixiong@users.noreply.github.com> --- nemo_rl/distributed/virtual_cluster.py | 4 +- nemo_rl/models/generation/trtllm/config.py | 80 +++++--- .../generation/trtllm/trtllm_disagg_server.py | 182 +++++++++--------- .../generation/trtllm/trtllm_generation.py | 96 +++++---- .../generation/trtllm/trtllm_http_server.py | 43 +++-- .../generation/trtllm/trtllm_worker_async.py | 4 +- .../trtllm/test_trtllm_disagg_server.py | 147 ++++++++++++++ .../trtllm/test_trtllm_generation.py | 105 ++++++++++ .../trtllm/test_trtllm_http_server.py | 2 + 9 files changed, 487 insertions(+), 176 deletions(-) diff --git a/nemo_rl/distributed/virtual_cluster.py b/nemo_rl/distributed/virtual_cluster.py index 7802aff3813..efb663f400b 100644 --- a/nemo_rl/distributed/virtual_cluster.py +++ b/nemo_rl/distributed/virtual_cluster.py @@ -1215,9 +1215,7 @@ def get_topology_sorted_pg_indices(self) -> list[int]: return identity topology = get_ray_cluster_topology() - if not any( - domain != NVLINK_DOMAIN_UNKNOWN for domain, _ in topology.values() - ): + if not any(domain != NVLINK_DOMAIN_UNKNOWN for domain, _ in topology.values()): return identity node_of_pg: list[str] = [] diff --git a/nemo_rl/models/generation/trtllm/config.py b/nemo_rl/models/generation/trtllm/config.py index 809a07339b1..f4f8658935e 100644 --- a/nemo_rl/models/generation/trtllm/config.py +++ b/nemo_rl/models/generation/trtllm/config.py @@ -12,12 +12,14 @@ # 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 TrtllmDisaggArgs(TypedDict): +class TrtllmDisaggConfig(BaseModel, extra="allow"): """Prefill/decode disaggregation. A *replica* is ``num_context_engines`` context engines plus @@ -31,11 +33,11 @@ class TrtllmDisaggArgs(TypedDict): together for the KV transceiver to work. """ - enabled: bool + enabled: bool = False # Engines per replica. The two are independent, so the P:D ratio is free. - num_context_engines: int - num_generation_engines: int + num_context_engines: int = 1 + num_generation_engines: int = 1 # Routing inside a replica, decided entirely by the disagg server. # @@ -48,44 +50,50 @@ class TrtllmDisaggArgs(TypedDict): # 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. - ctx_router: str # conversation | kv_cache_aware - gen_router: str # round_robin | load_balancing + # + # 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: NotRequired[int] + 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: NotRequired[bool] + 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: NotRequired[bool] + 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: NotRequired[bool] + frontend_tokenize: bool = False # Base port for the frontend workers' deterministic ports - # (base + frontend_idx on the pinned node). 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: NotRequired[int] + # (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. - # DEFAULT | UCX | NIXL | MOONCAKE | MPI - cache_transceiver_backend: str + cache_transceiver_backend: Literal["DEFAULT", "UCX", "NIXL", "MOONCAKE", "MPI"] = ( + "DEFAULT" + ) - # "CPP" | "PYTHON" | "auto". TRT-LLM defaults to "auto", which only adopts + # "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: NotRequired[str] + 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 @@ -93,23 +101,28 @@ class TrtllmDisaggArgs(TypedDict): # 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: NotRequired[int] - max_tokens_in_buffer: NotRequired[int] + 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: NotRequired[int] + 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. - ctx_trtllm_kwargs: NotRequired[dict[str, Any]] - gen_trtllm_kwargs: NotRequired[dict[str, Any]] + # 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): @@ -138,7 +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[TrtllmDisaggArgs] + disaggregation: NotRequired[TrtllmDisaggConfig] default_chat_template_kwargs: NotRequired[dict[str, Any]] # TRT-LLM's registered parser names: # "qwen3" -> Qwen3ToolParser (JSON format: {"name":..., "arguments":{...}}) @@ -153,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 index 25466ff21c0..3a091cf8e2f 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_disagg_server.py @@ -1,4 +1,3 @@ - # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); @@ -12,8 +11,7 @@ # 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``. +"""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 @@ -22,6 +20,7 @@ only hands it the two address pools and the router policies. """ +import json import logging import threading import time @@ -34,7 +33,9 @@ __all__ = [ "DisaggServerActor", "DisaggServerActorImpl", + "attach_rollout_fields", "build_config", + "generation_token_ids", "start_server", "wait_ready", ] @@ -111,6 +112,62 @@ def build_config( _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. @@ -133,9 +190,6 @@ async def __call__(self, scope: Any, receive: Any, send: Any) -> None: await self.app(scope, receive, send) return - - import json - chunks: list[bytes] = [] while True: message = await receive() @@ -315,66 +369,16 @@ async def wrapper(req: request_type, raw_req: Request) -> Response: # type: ign response = await inner(req, raw_req) if req.stream or not isinstance(response, JSONResponse): return response - return self._attach_rollout_fields(response, raw_req) - - return wrapper - - @staticmethod - 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( - self, response: JSONResponse, raw_req: "Request | None" = None - ) -> JSONResponse: - """Re-attach the fields NeMo-Gym reads off ``choices[].message``.""" - import json - - payload = json.loads(response.body) - - choices = payload.get("choices") or [] - if not choices: + # 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=payload, status_code=response.status_code + content=attach_rollout_fields(json.loads(response.body)), + status_code=response.status_code, ) - choice = choices[0] - message = choice.get("message") - if not isinstance(message, dict): - return JSONResponse( - content=payload, status_code=response.status_code - ) - - if payload.get("prompt_token_ids") is not None: - message["prompt_token_ids"] = payload["prompt_token_ids"] - - token_ids = self._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 JSONResponse(content=payload, status_code=response.status_code) + return wrapper return OpenAIDisaggServerAdaptor @@ -395,14 +399,13 @@ def _build_frontend_tokenizer( 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, ) - from tensorrt_llm.serve.openai_protocol import ChatCompletionRequest - tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True) model_config = AutoConfig.from_pretrained(model_name, trust_remote_code=True) @@ -575,31 +578,30 @@ def _start_serving( config, port=self._port, tokenize_fn=tokenize_fn ) - def start( - self, - ctx_addrs: Optional[list[tuple[str, int]]] = None, - gen_addrs: Optional[list[tuple[str, int]]] = None, - *, - ctx_router: str = "conversation", - gen_router: str = "load_balancing", - gen_tokids_ctxbytes: bool = False, - gen_strip_message_history: bool = False, - ) -> str: - """Ensure the server is up and return the URL NeMo-Gym will talk to.""" - if self._base_url is None: - assert ctx_addrs is not None and gen_addrs is not None, ( - "start() needs the address pools unless they were passed to " - "the constructor via serve_args" - ) - self._start_serving( - ctx_addrs, - gen_addrs, - ctx_router=ctx_router, - gen_router=gen_router, - gen_tokids_ctxbytes=gen_tokids_ctxbytes, - gen_strip_message_history=gen_strip_message_history, - ) + 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", @@ -628,8 +630,6 @@ def wait_ready(base_url: str, timeout_s: float = 300.0) -> None: It reaches out to every engine in its pools on startup, so readiness lags the thread start by more than a socket bind. """ - import time - import requests health = base_url.rsplit("/v1", 1)[0] + "/health" diff --git a/nemo_rl/models/generation/trtllm/trtllm_generation.py b/nemo_rl/models/generation/trtllm/trtllm_generation.py index 9a13985e3a0..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,13 +54,15 @@ def init_cluster_placement_groups( ) -> None: """Pre-initialize placement groups matching TRT-LLM's topology.""" trtllm_cfg = config["trtllm_cfg"] - disagg = trtllm_cfg.get("disaggregation") or {} + disagg = resolve_trtllm_disagg_config(config) engine_tp = trtllm_cfg["tensor_parallel_size"] - if disagg.get("enabled"): + if disagg.enabled: engine_tp = max( - int((disagg.get(f"{role}_trtllm_kwargs") or {}).get( - "tensor_parallel_size", engine_tp - )) + 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) @@ -83,7 +89,7 @@ def init_cluster_placement_groups( # 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.get("enabled") and colocated), ( + 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 " @@ -110,6 +116,10 @@ def __init__( ): self.cfg = config self.tp_size = self.cfg["trtllm_cfg"]["tensor_parallel_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 @@ -142,7 +152,7 @@ def __init__( # -- 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.get("enabled"): + if self._disagg_cfg.enabled: engine_configs = [ (f"{role}_trtllm_kwargs", self._role_kwargs(role)) for role in ("ctx", "gen") @@ -207,9 +217,7 @@ def __init__( # 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.get("enabled") - ) + 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 @@ -305,24 +313,24 @@ def __init__( # Engine layout # ------------------------------------------------------------------ # - def _role_kwargs(self, role: str) -> dict[str, Any]: + 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.get(f"{role}_trtllm_kwargs") or {} + 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]: - """Per-engine ``(role, tp_width)``, in engine order. + 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 ... + \______ 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 @@ -330,7 +338,7 @@ def _plan_engines(self, world_size: int) -> tuple[list[str], list[int], int]: are adjacent. """ disagg = self._disagg_cfg - if not disagg.get("enabled"): + 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})." @@ -346,8 +354,8 @@ def _plan_engines(self, world_size: int) -> tuple[list[str], list[int], int]: 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 = int(disagg["num_context_engines"]) - num_gen = int(disagg["num_generation_engines"]) + 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}" @@ -380,7 +388,7 @@ def _config_with_engine_overrides( 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.get("enabled"): + if node_bundle_indices is None or not self._disagg_cfg.enabled: return self.cfg overrides: dict[str, dict[str, Any]] = {} @@ -388,7 +396,7 @@ def _config_with_engine_overrides( for (pg_idx, bundles), role in zip( node_bundle_indices, self._engine_roles, strict=True ): - prefix = "ctx" if role == "context" else "gen" + 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 -- @@ -401,7 +409,7 @@ def _config_with_engine_overrides( overrides[self._engine_key(pg_idx, bundles)] = { "_disagg_role": role, "_disagg_role_ordinal": ordinal, - **(self._disagg_cfg.get(f"{prefix}_trtllm_kwargs") or {}), + **self._disagg_cfg.role_trtllm_kwargs(prefix), } cfg = dict(self.cfg) @@ -538,8 +546,8 @@ def _get_tied_worker_bundle_indices( # ------------------------------------------------------------------ # @property - def _disagg_cfg(self) -> dict[str, Any]: - return self.cfg["trtllm_cfg"].get("disaggregation") or {} + 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. @@ -550,7 +558,7 @@ def _assert_direct_dispatch_allowed(self) -> None: run *without* disaggregation while the user believes they are exercising it. Fail loudly instead. """ - if self._disagg_cfg.get("enabled"): + if self._disagg_cfg.enabled: raise RuntimeError( "PD disaggregation is only wired for the HTTP/NeMo-Gym rollout " "path; TrtllmGeneration.generate()/generate_async() dispatch " @@ -574,8 +582,8 @@ def _start_disagg_servers(self) -> list[Optional[str]]: return self._disagg_server_urls disagg = self._disagg_cfg - num_ctx = int(disagg["num_context_engines"]) - num_gen = int(disagg["num_generation_engines"]) + 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"), ( @@ -583,9 +591,12 @@ def _start_disagg_servers(self) -> list[Optional[str]]: "address is the only way the disagg server can reach an engine." ) - addrs = self._report_engine_addrs() - missing = [i for i, a in enumerate(addrs) if not a] + 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 @@ -593,7 +604,7 @@ def _start_disagg_servers(self) -> list[Optional[str]]: DisaggServerActor, ) - n_fe = int(disagg.get("num_frontend_workers") or 1) + 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" @@ -601,8 +612,13 @@ def _start_disagg_servers(self) -> list[Optional[str]]: # 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); frontends on one node get base+frontend_idx. - base_port = int(disagg.get("frontend_base_port") or 17300) + # (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 = [] @@ -616,13 +632,11 @@ def _start_disagg_servers(self) -> list[Optional[str]]: (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=bool(disagg.get("gen_tokids_ctxbytes", False)), - gen_strip_message_history=bool( - disagg.get("gen_strip_message_history", False) - ), - frontend_tokenize=bool(disagg.get("frontend_tokenize", False)), + 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" @@ -654,7 +668,7 @@ def _start_disagg_servers(self) -> list[Optional[str]]: replica_idx, frontend_idx=fe_idx, num_frontends=n_fe, - port=base_port + fe_idx, + port=base_port + replica_idx * n_fe + fe_idx, serve_args=serve_args, ) self._disagg_actors.append(actor) @@ -723,7 +737,7 @@ def _report_dp_openai_server_base_urls(self) -> list[Optional[str]]: 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.get("enabled"): + if self._disagg_cfg.enabled: return self._start_disagg_servers() if not self.cfg["trtllm_cfg"].get("expose_http_server"): return [cast(Optional[str], None)] * self.dp_size diff --git a/nemo_rl/models/generation/trtllm/trtllm_http_server.py b/nemo_rl/models/generation/trtllm/trtllm_http_server.py index 744c23fb38d..c39197a6199 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_http_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_http_server.py @@ -27,6 +27,7 @@ """ import asyncio +import functools import logging import os import random @@ -45,6 +46,21 @@ 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) @@ -306,8 +322,9 @@ async def chat_completions(request: Request): # 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. - _conv_id = ((body.get("conversation_params") or {}).get("conversation_id") - or (body.get("disaggregated_params") or {}).get("conversation_id")) + _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 ( @@ -318,9 +335,7 @@ async def chat_completions(request: Request): disagg_params = to_llm_disaggregated_params( WireDisaggregatedParams(**body["disaggregated_params"]) ) - is_context_leg = ( - getattr(disagg_params, "request_type", None) == "context_only" - ) + 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. @@ -474,9 +489,7 @@ async def chat_completions(request: Request): # -- 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 - ) + return _context_leg_response(model_name, adj_prompt, gen, disagg_params) gen_token_ids = list(gen.token_ids) @@ -562,10 +575,6 @@ async def chat_completions(request: Request): "index": 0, "message": msg_dict, "finish_reason": finish_reason, - # Generated token ids. ChatCompletionResponseChoice needs - # the matching field upstream (CompletionResponseChoice - # already has it) or the disagg server rejects this. - "token_ids": gen_token_ids, } ], # Declared on ChatCompletionResponse precisely so a generation @@ -578,6 +587,16 @@ 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 diff --git a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py index b6dacb50b79..b964f58d5a3 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py +++ b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py @@ -127,9 +127,7 @@ def _engine_overrides(self) -> dict[str, Any]: 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 - ) + 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( diff --git a/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py b/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py index 910d6d2c22d..afcd2a0aa9a 100644 --- a/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py +++ b/tests/unit/models/generation/trtllm/test_trtllm_disagg_server.py @@ -23,11 +23,17 @@ 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.""" @@ -133,3 +139,144 @@ def test_non_adapted_paths_pass_through_untouched(): 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): From f0a0d36bdac2eb2653e8c69009b3d1754e852bf4 Mon Sep 17 00:00:00 2001 From: shuyixiong <219646547+shuyixiong@users.noreply.github.com> Date: Tue, 22 Sep 2026 00:38:26 -0700 Subject: [PATCH 09/12] docs(trtllm): document every disaggregation field in the exemplar Nine of TrtllmDisaggConfig's fields appeared in no example config, so the only way to discover num_frontend_workers, frontend_base_port, gen_tokids_ctxbytes, gen_strip_message_history, frontend_tokenize, cache_transceiver_runtime, kv_cache_bounce_size_mb, max_tokens_in_buffer or kv_transfer_timeout_ms was to read the schema. Every value written here is the schema default, so the exemplar documents without becoming a second source of truth that can go stale against the BaseModel. The four Optional knobs are null, which is what "don't forward, let TRT-LLM keep its own default" actually looks like -- not the illustrative values a reader would otherwise mistake for defaults. ctx_trtllm_kwargs/gen_trtllm_kwargs become explicit empty dicts, and the comment now says a key set there WINS over trtllm_cfg: pinning tensor_parallel_size in the per-role block makes trtllm_cfg.tensor_parallel_size a silent no-op for the disagg layout, which is not obvious from the merge order. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: shuyixiong <219646547+shuyixiong@users.noreply.github.com> --- examples/configs/grpo_math_1B_trtllm.yaml | 47 ++++++++++++++++++++--- 1 file changed, 42 insertions(+), 5 deletions(-) diff --git a/examples/configs/grpo_math_1B_trtllm.yaml b/examples/configs/grpo_math_1B_trtllm.yaml index cf28abd4565..cbb82101916 100644 --- a/examples/configs/grpo_math_1B_trtllm.yaml +++ b/examples/configs/grpo_math_1B_trtllm.yaml @@ -36,12 +36,49 @@ policy: # 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). Uncomment to give the roles different shapes. - # ctx_trtllm_kwargs: - # tensor_parallel_size: 4 - # gen_trtllm_kwargs: - # tensor_parallel_size: 2 + # 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 From c168ca2553b548d8b7f2f0e89c777461194989a3 Mon Sep 17 00:00:00 2001 From: shuyixiong <219646547+shuyixiong@users.noreply.github.com> Date: Tue, 22 Sep 2026 02:01:11 -0700 Subject: [PATCH 10/12] docs(trtllm): record the Gym conversation-id dependency trtllm_http_server reads a conversation id that nothing populates today. Gym is the only component that knows the rollout identity and no released Gym sends it -- conversation_params appears nowhere in the tree, including upstream main. It arrives with NVIDIA-NeMo/Gym#3582 ("forward gym session id as backend conversation id"), still open, 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: 17-26% of turns find their prefix on the serving rank versus 96-97% once the id flows, measured at 2P-DEP8 / conc 512. Nothing fails, so the note also says not to treat such a run as a disaggregation baseline. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: shuyixiong <219646547+shuyixiong@users.noreply.github.com> --- .../generation/trtllm/trtllm_http_server.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/nemo_rl/models/generation/trtllm/trtllm_http_server.py b/nemo_rl/models/generation/trtllm/trtllm_http_server.py index c39197a6199..fde2c2840bb 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_http_server.py +++ b/nemo_rl/models/generation/trtllm/trtllm_http_server.py @@ -322,6 +322,21 @@ async def chat_completions(request: Request): # 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") From 74d75451bd7815774f7ff1d30918172c9e972394 Mon Sep 17 00:00:00 2001 From: shuyixiong <219646547+shuyixiong@users.noreply.github.com> Date: Tue, 22 Sep 2026 02:01:11 -0700 Subject: [PATCH 11/12] chore(trtllm): type-check trtllm_disagg_server.py pyrefly checks an explicit allow-list, and the new module was missing from it while its three siblings in the same package were listed -- so its 654 lines reported zero errors because nothing looked at them. It is clean once added. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: shuyixiong <219646547+shuyixiong@users.noreply.github.com> --- pyrefly.toml | 1 + 1 file changed, 1 insertion(+) 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", From 6cfd6a667a81e13b217685733a9cff505487fc6e Mon Sep 17 00:00:00 2001 From: shuyixiong <219646547+shuyixiong@users.noreply.github.com> Date: Tue, 22 Sep 2026 02:03:58 -0700 Subject: [PATCH 12/12] docs(trtllm): add the PD disaggregation rollout design doc Covers what the schema and the exemplar cannot: that NeMo RL owns the replica and the refit safety while delegating the KV handshake and the intra-replica routing to OpenAIDisaggServer, and that the replica count is derived from the inference cluster size rather than configured. A replica is fronted by num_frontend_workers disagg servers, not one, so the topology section is written in terms of the N*F URLs that reach dp_openai_server_base_urls and the three levels at which a trajectory acquires affinity: Gym's sticky sha256(session_id) picks the frontend, that frontend's in-process ctx_router picks the context engine, and the relayed conversation_id picks the attention-DP rank. Only the first is NeMo RL's to control, and the stickiness is a correctness requirement rather than load-spreading -- two frontends of one replica share no view of which context engine holds a conversation's prefix. Also records why the frontend's port carries both indices, why serving starts in the constructor (Ray replays __init__ on actor restart but never method calls), and that frontends are CPU-only so num_frontend_workers never enters the GPU budget. Listed in docs/index.md under Design Docs so Sphinx does not treat it as an orphan page. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: shuyixiong <219646547+shuyixiong@users.noreply.github.com> --- docs/design-docs/pd-disaggregation-rollout.md | 233 ++++++++++++++++++ docs/index.md | 1 + 2 files changed, 234 insertions(+) create mode 100644 docs/design-docs/pd-disaggregation-rollout.md 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