Skip to content

[RFC] Training/Inference Numerical Consistency for RL Post-Training #19501

Description

@shikicloud

Motivation

In RL post-training, the rollout engine samples tokens and the trainer recomputes logprobs on those same tokens with the same weights. When the two disagree, PPO's importance ratio is computed against the wrong denominator; low-probability tokens can spike the ratio and blow up gradients; training can collapse.

Four causes, each with published evidence:

  1. Kernels are not batch-invariant. The same request can produce different logits depending on what else is in the batch, because RMSNorm/matmul/attention pick their reduction strategy from batch shape (Thinking Machines, "Defeating Nondeterminism in LLM Inference", 2025-09-10).
  2. MoE routing disagrees between engines even with identical weights — ~10% of per-layer routing decisions differ, 94% of tokens affected in ≥1 layer, on Qwen3-30B-A3B (R3 paper, arXiv 2510.11370). R3 ([None][feat] Router Replay (R3): return per-token MoE routing to training engine in post train. #18397) is an initial solution to fix this cause.
  3. LM-head / router precision. BF16 alone dropped train/inference token-probability correlation to ~0.9; FP32 restored ~0.99 (MiniMax-M1 tech report, arXiv 2506.13585).
  4. Severity is hardware- and kernel-dependent and compounds with length. Same code: KL ≈ 5e-4–1e-3 on one GPU, 1e-2–1e-1 on another, traced to a specific attention-kernel bug (ByteDance Seed, "When Speed Kills Stability").

R3 addresses cause 2. This RFC covers what's left of R3 (cause 2's remaining edges) plus causes 1, 3, and a sampling-semantics issue that compounds all of them. Other useful features like TIS/MIS/KPop remain in RL frameworks but not inference engines. Futhermore, this RFC targets on important features for TRTLLM, we will explore new features to solve this problem on an new issue.

This RFC is generated by codex with human's instructions and supervision. Also thanks @NolenLiang and @JeffPengCoder's help. CC's @joyang-nv @xuantengh @yuki-97

What exists today

  • Router Replay: LlmArgs.enable_return_routed_experts + SamplingParams.return_routed_experts → CompletionOutput.routed_experts, int16, [seq_len-1, num_moe_layers, top_k], pre-EPLB logical ids. Separated-routing MoE backends only; fails closed on PP>1, spec decode, fused routing(TBD).
  • logprobs_mode: RAW (pre-sampling-params) / PROCESSED (post). No raw_logits/processed_logits variants.
  • FORCE_DETERMINISTIC: a C++ env var, not exposed via LlmArgs. Covers attention multi-block, KV-cache reuse, and all-reduce workspace size. Does not cover autotuner tactic selection, FMHA library selection, all-reduce strategy selection, or CUDA-graph padding(TBD).
  • MoeConfig.disable_finalize_fusion: CUTLASS only; the LlmArgs docstring states the fused path is non-deterministic for top-k > 2.

Two items already landed for Router Replay: per-request gating (SamplingParams.return_routed_experts=False now correctly attaches nothing), and the output dtype is int16.

Reading the PyTorch backend (tensorrt_llm/_torch/) turns up nine places where a numeric result depends on batch shape or token count, not just the sequence — this is what drives the Deterministic inference work below:

# Source Class Covered by FORCE_DETERMINISTIC?
1 Attention multi-block/split-KV (MMHA/XQA/TRTLLM-Gen); split count = f(SM_count, batch*heads) batch-dependent Yes — but a shared-memory guard can silently re-enable multi-block regardless of the flag
2 FMHA library selection by batch bucket (fmha/manager.py) batch-dependent No — TLLM_FMHA_LIBS can pin one manually
3 Dense GEMM tactic selection by M (autotuner) batch-dependent, plus profiling-noise without a persisted cache No — enable_autotuner=False falls back to tactic=-1, but the fallback cuBLASLt path's own algorithm choice is unexamined
4 MoE grouped-GEMM tactic selection by token count batch-dependent No
5 MoE finalize/unpermute reduction, CUTLASS top-k>2 non-deterministic per LlmArgs docstring Partially — disable_finalize_fusion=True; TRTLLM-Gen has no unfused path
6 All-reduce strategy selection by message size batch-dependent No — only workspace size is fixed, strategy selection isn't
7 CUDA-graph batch-size padding ladder batch-dependent No
8 Chunked-prefill / KV-reuse boundary placement depends on scheduler budget, not just the sequence Yes — both disabled

Two bugs found while reading the same code, independent of scope: getBoolEnv matches only the exact string "1" (FORCE_DETERMINISTIC=true is silently ignored), and MMHA's multi-block override under shared-memory pressure bypasses the flag entirely.

Proposed Change.

Organized by delivery phase. Each phase lists what happens to every workstream in that phase, rather than each workstream carrying its own separate phase numbering.

Phase 1 — config only, no new kernels, land first

  • Router Replay:
    • Validate TRTLLM-Gen, DeepGEMM, and WideEP end to end. Today only CUTLASS has a GPU test, on Qwen1.5-MoE-A2.7B-Chat; the other backends' support is inferred from _supports_load_balancer(), not tested.
    • Optional uint8 output for num_experts <= 256 models. int16 already halves the naive int32 footprint, but large payloads are a real failure mode at scale — a comparable rollout system hit ~35GB/actor, ~280GB across 8 actors, on a 96-layer MoE model with long context. Worth having the compact path ready before someone hits this on TRT-LLM.
  • Deterministic inference:
    • Config-only surface over the gaps in the table above:

      class TorchLlmArgs:
          enable_deterministic_inference: bool = False
          # Superset of FORCE_DETERMINISTIC=1:
          #   - autotuner forced off, one cached fallback tactic per (op, N, K)
          #     shape class instead of a per-M fallback
          #   - allreduce_strategy forced to one strategy (ONESHOT/MNNVL-oneshot)
          #     instead of AUTO's per-message-size lookup
          #   - one FMHA library pinned per (layer, phase)
          #   - CUDA-graph padding off, or tactics held constant across the ladder
          #   - disable_finalize_fusion forced on where supported
    • Fix the two bugs found above (getBoolEnv string match, MMHA shared-mem bypass) — independent of the rest of this phase, worth doing regardless.

    • Cost: comparable systems that shipped just this config layer report modest overhead since no kernels change; the bulk of the cost shows up in Phase 3.

  • Diagnostics:
    • A tests/unittest/_torch/determinism/ suite — same prompt at batch=1 vs. inside a large batch, assert bitwise-identical sampled-token logprobs, across attention and MoE backends. This is the highest-leverage single piece of infrastructure in the whole RFC: it validates Deterministic inference directly and is the tool used to sanity-check Router Replay and everything after it.
    • A weight_version counter on outputs, incremented on update_weights, so a mid-generation refit is attributable to the right requests.

Phase 2 — small API/kernel changes

  • LM head and router precision:
    • Problem: logits_processor.py computes the LM-head GEMM in BF16 and only casts the result to fp32 afterward — the GEMM's own accumulation stays BF16. LM Head fp32: [None][feat] Add LlmArgs.lm_head_dtype for float32 LM head logits #19850

    • Proposal:

      class TorchLlmArgs:
          lm_head_dtype: Optional[torch.dtype] = None
          # fp32 the LM-head GEMM itself, not just the output cast.
          # Unquantized lm_head only.
      
      class MoeConfig:
          router_dtype: Optional[torch.dtype] = None
          # Generalizes the existing DeepSeek-only fp32 gate to any MoE model;
          # fixes the Qwen3 FIXME as a side effect.
    • Also fixes an existing bug: Qwen3's MoE gate has a broken fp32 attempt already in the code (# FIXME: out_dtype=float32 does not work) — worth landing as a standalone patch regardless of the rest of this RFC.

    • Caveat: fp32 alone hasn't been sufficient to eliminate mismatch in other teams' reported experiments — it helped in cases where the bottleneck happened to be specifically there. Scope this as closing one known precision-sensitive path, not a general fix.

  • Sampling mask replay:
    • Problem: when top_p/top_k truncate the distribution, PROCESSED logprobs are renormalized over the truncated support. If the trainer normalizes π_θ over the full vocabulary instead, the resulting importance ratio is systematically biased — no amount of trainer-side ratio clipping/masking (TIS, MIS) can fix this, since it's a mismatched support, not a bad ratio.

    • Proposal:

      class TorchLlmArgs:
          enable_return_sampling_mask: bool = False
      
      class SamplingParams:
          return_sampling_mask: bool = False
          # Requires logprobs_mode=PROCESSED, temperature>0, finite top_k.
      
      class CompletionOutput:
          sampling_mask: Optional[list[list[int]]]
          # Per step, the token ids that survived filtering — the support the
          # trainer should renormalize over.
    • Reuses Router Replay's capture-buffer pattern (route_capture.py's device-buffer + async D2H) rather than inventing a new one.

  • Router Replay:
    • Skip re-returning already-known prefix routes on multi-turn requests. Not urgent until multi-turn RL usage on this engine materializes, but cheap to build alongside the sampling-mask work since it reuses the same buffer plumbing.
  • Diagnostics:

Phase 3 — kernel work

  • Deterministic inference:
    • Attention: fixed split-size, not fixed split-count. Fixing the count still shifts chunk boundaries as total KV length changes, so the reduction order for a given KV position keeps moving; fixing the size anchors it.
    • Dense GEMM: a fixed-tile, no-split-K path for small-M BF16 layers, to stop depending on cuBLASLt's internal shape-based heuristics.
    • TRTLLM-Gen MoE: needs either an unfused finalize path or a fixed-order fused one — it has neither today.
    • RMSNorm: audit the NVFP4 M >= 4096 kernel-switch gate (likely fine, unverified).
    • Cost: expect roughly 20–35% overhead once this lands, based on comparable systems' reported numbers for the equivalent kernel work.
  • Router Replay:
    • Speculative decoding / MTP support. Fails closed today — draft/verification tokens need routing captured too, or the feature silently stops working the moment spec decode is turned on for rollout throughput.
    • Pipeline parallelism, PP>1. Fails closed today; lower priority since PP is less commonly used on the rollout side of an RL setup.

Backlog / stretch

  • Deterministic inference:
    • TP-size invariance — fixing intra-GPU and inter-GPU reduction to the same order so a TP>1 rollout and a TP=1 trainer agree (TBIK, arXiv 2511.17826). Research-grade; not committed scope.
  • Router Replay:
    • KV-offload / PD-disaggregation compatibility. Fails closed on any KV connector today.
    • Streaming responses don't currently carry routed_experts. Needs an explicit decision either way.
    • Longer-term: the same record-and-replay idea generalizes to sparse-attention index selection (DSA-style top-k KV selection), if/when TRT-LLM ships that.
  • Weight sync fidelity:
    • Rollout Engine and Trainer use the exact same kernels.
  • Weight sync fidelity:
    • When rollout runs in a lower precision than the trainer's weights (FP8/NVFP4 rollout against BF16-trained weights), each refit (update_weights) has to requantize using exactly the scheme the serving numerics were validated against. A scale or rounding mismatch here opens a new, independent gap — one unrelated to Router Replay or Deterministic inference, and just as capable of causing training instability on its own (FlashRL documents this failure mode directly: uncorrected quantized rollout collapses entropy and requires an importance-sampling correction to recover BF16-level accuracy).
    • Out of scope for now: closer to a weight-loading/quantization problem than an inference-numerics one. The fix likely splits across the calling framework (get the quantization scheme right on its side) and TRT-LLM (validate and document the contract, and cover it in the diagnostics work above). Flagged here as a recognized gap, not a committed design.

Docs for each item above land alongside the phase that ships it — logprobs_mode semantics, Router Replay usage and limits, weight-sync (rlhf_utils.WorkerExtension), sleep/wake, and the new flags as they ship — rather than as one deferred deliverable at the end.

API surface

Field Type Phase
LlmArgs.enable_deterministic_inference bool 1
CompletionOutput.weight_version int 1
LlmArgs.routed_experts_dtype (uint8/int16) torch.dtype 1, optional
LlmArgs.lm_head_dtype Optional[torch.dtype] 2
MoeConfig.router_dtype Optional[torch.dtype] 2
LlmArgs.enable_return_sampling_mask / SamplingParams.return_sampling_mask bool 2
CompletionOutput.sampling_mask Optional[list[list[int]]] 2
SamplingParams.routed_experts_start_len int 2

References

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions