Rollout Routing Replay (R3) in Miles: Preventing MoE Routing Mismatch for Stable RL Training
Rollout Routing Replay (R3) is a mechanism in Miles that records expert routing decisions during inference and replays them during training to eliminate stochastic mismatches caused by FP8 quantization and non-deterministic kernels, ensuring bit-identical expert allocation and stable MoE reinforcement learning.
Mixture-of-Experts (MoE) models power modern large language models, but their inherent routing stochasticity creates a critical failure mode in RL training pipelines. The Miles training framework solves this with Rollout Routing Replay (R3)—a system that captures routing metadata during rollout and injects it back during the training forward pass. This article explains how R3 works, why MoE RL fails without it, and how to implement it in your training runs.
Why MoE RL Suffers From Routing Mismatch
In MoE architectures, a learned router (nn.Linear) assigns each token to its top-k most relevant experts. This process is sensitive to tiny numerical perturbations:
- FP8 quantization and non-deterministic GPU kernels introduce noise in the router's output logits
- The top-k selection operation amplifies small differences, potentially routing the same token to entirely different experts between rollout and training
- When training recomputes the forward pass with mismatched experts, gradients flow to the wrong parameters, causing policy divergence
As documented in docs/advanced/miles-router.md, without R3 these mismatches accumulate rapidly and destabilize training【/cache/repos/github.com/radixark/miles/main/docs/advanced/miles-router.md†L16-L23】.
How R3 Records and Replays Routing Decisions
The R3 mechanism operates across two phases with minimal overhead.
Rollout Phase: Capture Routing Metadata
When launching with --use-rollout-routing-replay, the SGLang inference engine extends each response's meta_info with a routed_experts tensor:
- Shape:
(seq_len-1, num_layers, top_k) - Dtype:
int32 - Content: Exact expert indices selected by the router for every token at every layer
This tensor is stored in sample.rollout_routed_experts, defined in miles/utils/types.py【/cache/repos/github.com/radixark/miles/main/miles/utils/types.py†L99-L102】.
Training Phase: Inject Recorded Routes
The RoutingReplayManager class in miles/utils/replay_base.py orchestrates replay:
# Core replay cycle as implemented in BaseReplayManager
from miles.utils.replay_base import RoutingReplayManager, Replay
# Record stage: capture router's top-k decisions
replay = Replay() # Stores indices in forward_list / backward_list
replay.record(indices) # Called during original forward pass
# Replay forward: bypass router, use stored indices
indices = replay.pop_forward() # Returns recorded expert assignments
# Replay backward: ensure gradient consistency
indices = replay.pop_backward() # Same indices for backward computation
The manager supports two operational modes:
replay_forward: Uses stored indices, skips router computation entirelyreplay_backward: Supplies identical indices for gradient calculation
Optional Safety Validation
For debugging, enable consistency checking via BaseReplayManager.check_replay_result【/cache/repos/github.com/radixark/miles/main/miles/utils/replay_base.py†L58-L66】:
routing_replay_manager.enable_check_replay_result = True
routing_replay_manager.replay_check_min_overlap_ratio = 0.95 # 95% match required
This compares replayed indices against freshly computed routes and raises an error if divergence exceeds the threshold.
Memory and Performance Characteristics
R3 adds modest storage overhead for routing metadata:
| Configuration | Memory Per Sequence |
|---|---|
| 32K tokens, 60 layers, top_k=8 | ~60 MB |
| 16K tokens, 60 layers, top_k=8 | ~30 MB |
R3 automatically disables for dense models and certain advantage estimators that already mask off-policy terms, avoiding unnecessary overhead【/cache/repos/github.com/radixark/miles/main/docs/advanced/miles-router.md†L53-L58】.
Enabling R3 in Your Training Run
Launch-Time Configuration
# Add to any Miles run script
python scripts/run_qwen3_dense.py \
--use-rollout-routing-replay \
# ... other arguments
Training Loop Integration
from miles.utils.replay_base import routing_replay_manager
# Activate replay for this batch
routing_replay_manager.enabled = True
routing_replay_manager.stage = "replay_forward"
# Model automatically uses recorded routing—no additional code required
output = model(sample.tokens)
loss = compute_loss(output, sample.labels)
loss.backward()
# Replay manager handles backward consistency automatically
Debug Mode With Validation
# Enable strict consistency checking
routing_replay_manager.enable_check_replay_result = True
routing_replay_manager.replay_check_min_overlap_ratio = 0.90 # ≥90% overlap required
Key Implementation Files
| File | Purpose |
|---|---|
docs/advanced/miles-router.md |
High-level R3 documentation and configuration guide【/cache/repos/github.com/radixark/miles/main/docs/advanced/miles-router.md†L30-L34】 |
miles/utils/types.py |
Sample.rollout_routed_experts tensor definition【/cache/repos/github.com/radixark/miles/main/miles/utils/types.py†L99-L102】 |
miles/utils/replay_base.py |
RoutingReplayManager, BaseReplayManager, and Replay classes【/cache/repos/github.com/radixark/miles/main/miles/utils/replay_base.py†L19-L26】 |
miles/utils/replay_base.py (check_replay_result) |
Optional consistency validation logic【/cache/repos/github.com/radixark/miles/main/miles/utils/replay_base.py†L58-L66】 |
scripts/run_qwen3_dense.py |
Example launch script with --use-rollout-routing-replay flag |
Summary
- Rollout Routing Replay (R3) eliminates MoE routing mismatch by capturing expert assignments during inference and replaying them during training
- The mechanism guarantees bit-identical expert allocation, preventing gradient corruption from FP8 and kernel non-determinism
- Core components:
RoutingReplayManager,Replaystorage class, androllout_routed_expertstensor inSampleobjects - Memory overhead is ~60 MB per 32K-token sequence; R3 auto-disables for inappropriate configurations
- Enable with
--use-rollout-routing-replayflag and optional consistency checking for production training runs
Frequently Asked Questions
What causes MoE routing mismatch without R3?
Non-deterministic GPU kernels, FP8 quantization noise, and floating-point precision differences between rollout and training phases cause the router to select different experts for identical tokens. The top-k operation amplifies these small numerical differences into entirely different expert assignments, corrupting gradients and destabilizing policy training.
How much memory does R3 consume?
R3 stores routing metadata as int32 tensors with shape (seq_len-1, num_layers, top_k). For a 32,000-token sequence with 60 layers and top_k=8, this requires approximately 60 MB per sample. The mechanism automatically disables for dense models and certain advantage estimators to avoid unnecessary overhead.
Can I verify that R3 is working correctly?
Yes. Set routing_replay_manager.enable_check_replay_result = True and configure replay_check_min_overlap_ratio (default 0.95) to validate that replayed indices match freshly computed routes within your tolerance. The manager raises an error if divergence exceeds the threshold, helping catch implementation bugs or unexpected numerical behavior.
When should I disable R3?
R3 disables automatically for dense models and specific advantage estimators like GRPO that already mask off-policy terms. You should not manually disable R3 for MoE RL training—doing so risks the routing mismatch failures it was designed to prevent.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →