Get Started
Home
Topics
Search
Library
6 min read · Inference Optimization · LLM Training · Added Oct 8 · Paper published Oct 5, 2026

TRIAGE: Direction-Aware Mismatch Stabilization of Native NVFP4 Reinforcement Learning

Source: research paper via Hugging Face Daily Papers
4-bit RL training collapses because learner and sampler kernels disagree on token probabilities, breaking GRPO’s importance weighting. TRIAGE reframes the fix: instead of clipping large mismatches, down-weight only tokens whose sign makes the gap self-amplify. Result: 2.3× rollout throughput on Blackwell, matching BF16 accuracy.
TL;DR
TRIAGE stabilizes 4-bit NVFP4 reinforcement learning by distinguishing policy-gradient updates that amplify learner-sampler probability mismatch from those that contract it, selectively down-weighting the amplifying ones, which lets native W4A4 execution reach ~2.3× the rollout throughput of BF16 without the usual training collapse.
Why It Matters
When you RL-finetune a reasoning LLM, generating the rollouts (long chains of thought) dominates wall-clock time. NVIDIA’s Blackwell GPUs can run matrix multiplies in a 4-bit format called NVFP4 at roughly 4× BF16 throughput, so a natural play is: sample rollouts in 4-bit, then compute gradients in 4-bit too.
The problem: in modern RL stacks the sampler (e.g. SGLang or vLLM) and the learner (e.g. Megatron-LM) are two separate execution paths. Even with identical weights, 4-bit kernels produce slightly different token probabilities than higher-precision ones. That gap, called learner-sampler mismatch, breaks the importance-weighting assumption behind Group Relative Policy Optimization (GRPO) and often causes training to diverge after a few hundred steps.
Prior fixes attack the magnitude of this mismatch: truncated importance sampling (TIS (Truncated Importance Sampling)), clipping, or masking tokens whose learner/sampler probability ratio is too far from 1. The authors argue this misses the point. A big mismatch may already be self-correcting under the next gradient step, while a small one can systematically grow.
How It Works
The core insight is a sign-based decomposition of what the next policy-gradient step does to the mismatch. For each token, let δ be the learner-minus-sampler log-probability gap, and let the token’s advantage A carry the sign of the update (positive advantage pushes the probability up, negative pushes it down).
A first-order expansion of the squared-mismatch “energy” shows that each token’s contribution to mismatch change has the sign of A·δ. So two regions are self-amplifying:
•
A<0, δ<0: learner already underweights the token, and the update pushes it lower.
•
A>0, δ>0: learner already overweights it, and the update pushes it higher.
The other two sign combinations are self-contracting. In their NVFP4 runs, the authors observe that the A<0, δ<0 region fills up disproportionately before any broad explosion of |δ|, and that these dangerous tokens cluster in a few 64-token segments inside long responses. Response-level averages hide them: at the pre-terminal stage, ~97% of affected 4B responses still have mean |δ| below 0.05.
TRIAGE acts on this with two pieces, both leaving the forward pass in native W4A4:
•
Direction-selective gating. Chop each response into 64-token segments. If a segment has a persistently negative mean gap or a heavy negative tail, attenuate only the amplifying tokens inside it (A<0, δ<0), not the whole segment. Naturally-contracting tokens in the same segment keep full weight.
•
Bounded repair. On positive-advantage responses whose segment mean gap drifts too negative, add a Pseudo-Huber penalty penalty that nudges learner log-probabilities back up. Repair kicks in only after a fixed update step and is bounded so it cannot dominate the policy gradient.
for segment S in response i: mean_gap, tail_frac = stats(S) w_S = gate(mean_gap, tail_frac) # in [w_min, 1] for token t in S: if A[i] < 0 and delta[t] < 0: w[t] = w_S # attenuate amplifiers else: w[t] = 1 loss_pg = weighted_response_mean(w, base_loss) if step >= k_rep: loss += lambda_rep * pseudo_huber(positive_adv_segment_gaps)
What They Found
Experiments use Qwen3-4B-Base and the MoE model Qwen3-30B-A3B-Base, trained with Group Relative Policy Optimization (GRPO) on the DAPO math dataset, with max response length 20,480 tokens.
•
Stability. Naive NVFP4 collapses after ~300 steps on 4B and shows runaway mismatch after ~700 steps on 30B. NVFP4+TIS (Truncated Importance Sampling) delays but does not prevent 4B collapse. NVFP4+TRIAGE stays stable across the full 600-step (4B) and 1,700-step (30B) horizons.
•
Benchmark accuracy. Averaged across GSM8K, MATH-500, AIME 2024/2025, and AMC 2023, TRIAGE reaches 58.49% on 4B vs 58.26% for BF16, and 70.96% on 30B vs 72.41% for BF16. At the shared step-300 checkpoint on 4B, TRIAGE beats naive NVFP4 by 5.93 pp and NVFP4+TIS by 2.27 pp, so the gain is not just from training longer.
•
Throughput. On eight B300 GPUs, native NVFP4 with TRIAGE delivers up to 2.30× BF16 rollout throughput and 1.30× end-to-end iteration speedup. The TRIAGE-specific learner overhead is a median 0.84%.
•
Ablation. Removing repair keeps mean mismatch low but leaves the directional imbalance uncorrected and loses reward. TIS alone, started from a healthy TRIAGE checkpoint, stays stable, which the authors read as evidence that the critical intervention is early-stage direction control, not late-stage magnitude clipping.
One caveat the authors flag themselves: on the 30B model, the measured gap is a residual after replaying MoE routing, and their per-token audit of the “amplifying = sign(A·δ)” rule only holds up as a token-level predictor on the 4B dense model. On 30B they only claim the population-level pattern.
What’s Useful
•
If you run RL on Blackwell-class GPUs and rollout is your bottleneck, this paper is a concrete argument that native W4A4 is viable, provided you intervene on the optimization objective rather than on the quantization recipe. Rollout throughput roughly doubles; end-to-end speedup is more modest (~1.3×) because parameter sync and quantization reductions eat the rest.
•
If you already use TIS (Truncated Importance Sampling) or clipping and still see slow drift or late-stage reward decay, it’s worth testing whether your failures correlate with the A<0, δ<0 region concentrating in a few segments. The paper’s diagnostic (segment-mean gap plus a tail-fraction count over 64-token windows) is cheap to add as a monitor even without adopting the full method.
•
The directional decomposition is independent of 4-bit execution. Asynchronous RL, where policy staleness creates a similar learner-sampler gap, is called out as a plausible next application, but the authors have not tested it.
•
Do not read the result as “4-bit RL is solved.” The evaluated models top out at 30B active-expert MoE and training runs stop at convergence or stable late behavior, not at the 100B+ scale.
Caveats
•
All evidence is math-reasoning benchmarks with GRPO; no coding, agentic, or general-purpose RL tasks.
•
The 30B measurements use routing-replay, so its absolute mismatch numbers are not directly comparable to the 4B dense run, and the per-token directional predictor is weaker there.
•
End-to-end speedup is bounded by non-GEMM overhead (parameter sync, quantization scale computation); the headline 2.3× is rollout-only.
•
Gate thresholds, segment length (64 tokens), and the repair activation step are hand-tuned and fixed per run; the paper does not study sensitivity to these choices beyond a W=128 segment-length check.
Topics
Inference Optimization
LLM Training
Reinforcement Learning
Inference Optimization
LLM Training
Reinforcement Learning
Up next in Inference Optimization
NeMo-DCR: Bit-Exact Delta-Compressed Refit for Scalable Agentic RL at Trillion-Parameter Scale
SlimWise: Decoupling Expert Pruning Across Prefill and Decode for Efficient MoE Serving
Don't miss new content
Log in to follow topics and personalize your feed.
Related topics you might like
Reinforcement Learning119 episodes
Inference Optimization135 episodes
LLM Training164 episodes