Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 30 additions & 2 deletions nemo_rl/models/generation/megatron/megatron_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,14 @@ def _initialize_inference_engine(self, mcore_generation_config: dict) -> None:
)
from megatron.core.utils import get_attr_wrapped_model

from nemo_rl.models.megatron.router_replay import (
assert_megatron_inference_router_replay_ready,
rebuild_global_router_replay_registry,
reset_moe_routing_metadata_buffer,
router_replay_enabled,
sync_model_config_for_router_replay,
)

pg_collection = get_attr_wrapped_model(self.model, "pg_collection")

buffer_size_gb = mcore_generation_config["buffer_size_gb"]
Expand All @@ -107,7 +115,10 @@ def _initialize_inference_engine(self, mcore_generation_config: dict) -> None:

# The value may be overwritten by `recompute_kv_cache_after_weight_updates`.
kv_cache_management_mode = mcore_generation_config["kv_cache_management_mode"]
needs_static_kv_pointers = kv_cache_management_mode != "persist"
static_kv_memory_pointers = mcore_generation_config.get(
"static_kv_memory_pointers",
kv_cache_management_mode != "persist",
)

materialize_only_last_token_logits = mcore_generation_config[
"materialize_only_last_token_logits"
Expand All @@ -117,6 +128,14 @@ def _initialize_inference_engine(self, mcore_generation_config: dict) -> None:

mamba_inference_state_config = MambaInferenceStateConfig.from_model(self.model)
is_hybrid_model = mamba_inference_state_config is not None
if is_hybrid_model and self.cfg.get("megatron_cfg", {}).get(
"zero_train_gen_mismatch"
):
# Match the train scan's fp32 boundary-state precision so gen SSM cache
# doesn't diverge from train (gen defaults to bf16 otherwise).
mcore_generation_config.setdefault(
"mamba_inference_ssm_states_dtype", "float32"
)
if is_hybrid_model:
if (
mcore_generation_config.get("mamba_inference_ssm_states_dtype")
Expand All @@ -139,6 +158,10 @@ def _initialize_inference_engine(self, mcore_generation_config: dict) -> None:
if logging_step_interval is None:
logging_step_interval = 0

if router_replay_enabled(self.cfg):
rebuild_global_router_replay_registry(self.model)
sync_model_config_for_router_replay(self.model, self.cfg)

# flashinfer's fused-RoPE kernel only dispatches fp16/bf16 q/k.
use_flashinfer_fused_rope = self.model.config.params_dtype in (
torch.float16,
Expand All @@ -152,7 +175,7 @@ def _initialize_inference_engine(self, mcore_generation_config: dict) -> None:
max_tokens=max_tokens,
max_sequence_length=mcore_generation_config["max_model_len"],
kv_cache_management_mode=KVCacheManagementMode(kv_cache_management_mode),
static_kv_memory_pointers=needs_static_kv_pointers,
static_kv_memory_pointers=static_kv_memory_pointers,
use_cuda_graphs_for_non_decode_steps=use_cuda_graphs_for_non_decode_steps,
use_flashinfer_fused_rope=use_flashinfer_fused_rope,
sampling_backend="flashinfer",
Expand Down Expand Up @@ -182,6 +205,11 @@ def _initialize_inference_engine(self, mcore_generation_config: dict) -> None:
self.inference_context = DynamicInferenceContext(
self.model.config, inference_config
)
if router_replay_enabled(self.cfg):
reset_moe_routing_metadata_buffer(self.inference_context)
assert_megatron_inference_router_replay_ready(
self.model, self.inference_context, self.cfg
)
self.inference_wrapped_model = GPTInferenceWrapper(
self.model, self.inference_context
)
Expand Down
96 changes: 96 additions & 0 deletions nemo_rl/models/megatron/moe_routing_record.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
# 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.
"""Record MoE routing during Megatron-inference generation as R3 ``routed_experts``.

Upstream router replay (R3, ``nemo_rl.models.megatron.router_replay``) provides the
*replay* side for the Megatron training/logprob forward, but only records routing on
the vLLM generation path. The colocated Megatron-inference generation path therefore
has no recorder. This module bridges that gap: it converts the dynamic-inference
engine's per-sample ``routing_indices`` into the ``[batch, seq, layers, topk]``
``routed_experts`` tensor that ``router_replay.set_router_replay_forward`` /
``build_router_replay_assignments`` consume (see ``_normalize_routed_experts_for_mcore``,
which accepts ``[B, S, L, K]``).
"""

from typing import Optional

import torch


def coerce_routing_to_3d(routing: torch.Tensor) -> torch.Tensor:
"""Normalize per-sample routing to ``[num_tokens, num_layers, topk]``."""
if not isinstance(routing, torch.Tensor):
raise TypeError(f"routing_indices must be a torch.Tensor, got {type(routing)}")
if routing.ndim == 3:
return routing
raise ValueError(
f"routing_indices must be 3D [tokens, layers, topk], got shape {tuple(routing.shape)}"
)


def align_routing_rows_to_token_count(
routing: torch.Tensor, num_tokens: int
) -> torch.Tensor:
"""Pad/trim routing rows to ``num_tokens`` so they replay on a full-sequence forward.

Dynamic inference accumulates ``[num_tokens, layers, topk]`` (sometimes one fewer
row than the padded sequence length). Repeat the last recorded row when an extra
step is required; trim when there are more.
"""
if routing.ndim != 3:
raise ValueError(f"Expected 3D routing tensor, got {tuple(routing.shape)}")
num_rows = routing.shape[0]
if num_rows == num_tokens:
return routing
if num_rows > num_tokens:
return routing[:num_tokens].contiguous()
if num_rows == 0:
raise ValueError("Cannot align empty routing indices to a non-empty sequence")
pad_rows = num_tokens - num_rows
last = routing[-1:].expand(pad_rows, -1, -1)
return torch.cat([routing, last], dim=0)


def build_routed_experts_batch(
routing_per_sample: list[Optional[torch.Tensor]],
seq_lengths: torch.Tensor,
seq_dim: int,
) -> Optional[torch.Tensor]:
"""Build the R3 ``routed_experts`` tensor ``[batch, seq_dim, layers, topk]``.

Each sample's routing is aligned to its (unpadded) sequence length and right-padded
with zeros to ``seq_dim`` (the generation ``output_ids`` sequence length), so that
``routed_experts[i][:seq_lengths[i]]`` is the per-token routing and the tail is pad.
Returns ``None`` when no sample carries routing (e.g. dense models / replay disabled).
"""
if not any(r is not None for r in routing_per_sample):
return None
if any(r is None for r in routing_per_sample):
raise ValueError(
"routing_indices must be present for every sample when router replay is enabled"
)
aligned = [
align_routing_rows_to_token_count(
coerce_routing_to_3d(r), int(seq_lengths[i].item())
)
for i, r in enumerate(routing_per_sample)
]
num_layers, topk = aligned[0].shape[1], aligned[0].shape[2]
out = torch.zeros(
len(aligned), seq_dim, num_layers, topk, dtype=aligned[0].dtype
)
for i, r in enumerate(aligned):
n = min(r.shape[0], seq_dim)
out[i, :n] = r[:n]
return out
119 changes: 117 additions & 2 deletions nemo_rl/models/megatron/router_replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,15 +47,70 @@ def configure_vllm_for_router_replay(config: PolicyConfig) -> None:
vllm_kwargs["enable_return_routed_experts"] = True


def sync_model_config_for_router_replay(model: Any, config: PolicyConfig) -> None:
"""Ensure the built mcore model config enables inference routing recording."""
if not router_replay_enabled(config):
return
model_config = _unwrap_model_config(model)
if model_config is None:
return
model_config.moe_enable_routing_replay = True
# Fused TE top-k bypasses RouterReplay.RECORD; force the replay-capable path.
model_config.moe_router_fusion = False


def assert_megatron_inference_router_replay_ready(
model: Any, inference_context: Any, config: PolicyConfig
) -> None:
"""Fail fast when router replay is on but the inference engine cannot record routing."""
if not router_replay_enabled(config):
return

from megatron.core.transformer.moe.router_replay import RouterReplay

model_config = getattr(model, "config", None)
if model_config is None or not getattr(
model_config, "moe_enable_routing_replay", False
):
raise RuntimeError(
"policy.router_replay.enabled=true requires model.config."
"moe_enable_routing_replay=True before Megatron inference starts."
)

if getattr(model_config, "num_moe_experts", None) in (None, 0):
raise RuntimeError(
"policy.router_replay.enabled=true requires a MoE model with "
"num_moe_experts > 0."
)

if not RouterReplay.global_router_replay_instances:
raise RuntimeError(
"policy.router_replay.enabled=true but no RouterReplay instances "
"were registered on the model. Ensure setup_model_config applied "
"moe_enable_routing_replay before the model was built."
)

if getattr(inference_context, "moe_routing_metadata", None) is None:
raise RuntimeError(
"policy.router_replay.enabled=true but the Megatron dynamic "
"inference context has no moe_routing_metadata. The inference "
"engine was likely initialized before moe_enable_routing_replay "
"was enabled; restart the job after enabling router replay."
)


def validate_router_replay_config(config: PolicyConfig) -> None:
if not router_replay_enabled(config):
return

generation = config.get("generation") or {}
megatron_cfg = config.get("megatron_cfg") or {}

if generation.get("backend") != "vllm":
raise ValueError("router_replay.enabled requires vLLM generation.")
generation_backend = generation.get("backend")
if generation_backend not in ("vllm", "megatron"):
raise ValueError(
"router_replay.enabled requires vLLM or Megatron generation."
)
if not megatron_cfg.get("enabled", False):
raise ValueError("router_replay.enabled requires the Megatron policy backend.")

Expand Down Expand Up @@ -99,8 +154,37 @@ def _unwrap_model_config(model: Any) -> Optional[Any]:
return None


def _hybrid_moe_layer_numbers(hybrid_layer_pattern: str, num_layers: int) -> list[int]:
# Deferred import: megatron.core is only available inside the Megatron worker venv.
from megatron.core.models.hybrid.hybrid_layer_allocation import (
Symbols,
parse_hybrid_pattern,
)

main_pattern = parse_hybrid_pattern(hybrid_layer_pattern).main_pattern or ""
# '|' marks a pipeline-stage boundary, not a layer; every other symbol occupies one slot.
layer_symbols = [symbol for symbol in main_pattern if symbol != Symbols.PIPE]
if len(layer_symbols) != num_layers:
raise ValueError(
f"hybrid_layer_pattern main segment has {len(layer_symbols)} layers "
f"but num_layers={num_layers} (pattern={hybrid_layer_pattern!r})"
)
return [
layer_idx + 1
for layer_idx, symbol in enumerate(layer_symbols)
if symbol == Symbols.MOE
]


def _global_moe_layer_numbers(model_config: Any) -> list[int]:
num_layers = int(getattr(model_config, "num_layers"))

# Hybrid Mamba/attention models place MoE layers via 'E' symbols of hybrid_layer_pattern
# and leave moe_layer_freq at its default of 1, which would wrongly claim every layer is MoE.
hybrid_layer_pattern = getattr(model_config, "hybrid_layer_pattern", None)
if hybrid_layer_pattern:
return _hybrid_moe_layer_numbers(hybrid_layer_pattern, num_layers)

moe_layer_freq = getattr(model_config, "moe_layer_freq", 1)

if isinstance(moe_layer_freq, int):
Expand Down Expand Up @@ -527,3 +611,34 @@ def clear_global_router_replay_instances() -> None:
from megatron.core.transformer.moe.router_replay import RouterReplay

RouterReplay.clear_global_router_replay_instances()


def rebuild_global_router_replay_registry(model: Any) -> None:
"""Re-register policy RouterReplay instances after a temporary model build.

Reference-model setup builds a throwaway Megatron model whose RouterReplay
objects register globally, then clears the global list. Megatron inference
recording uses that global list, so we must restore the live policy model's
instances before generation starts.
"""
from megatron.core.transformer.moe.router_replay import RouterReplay
from megatron.core.utils import unwrap_model

instances = _router_replay_instances_for_model(unwrap_model(model))
if not instances:
return

RouterReplay.clear_global_router_replay_instances()
# Preserve module-walk order to match RouterReplay() instantiation order.
RouterReplay.global_router_replay_instances.extend(
replay_instance for replay_instance, _ in instances
)


def reset_moe_routing_metadata_buffer(inference_context: Any) -> None:
"""Drop cached MoE routing CUDA-graph buffers so they resize to the live registry."""
metadata = getattr(inference_context, "moe_routing_metadata", None)
if metadata is None:
return
metadata.routing_indices_buffer = None
metadata.num_moe_layers = None
Loading
Loading