From 9d4dd040ae12f8c54a4b4dc82a85bef47d646346 Mon Sep 17 00:00:00 2001 From: Jorge Albericio Date: Tue, 6 Oct 2026 11:17:32 -0400 Subject: [PATCH] feat(hybrid): add opt-in shared-prefix attention and Mamba execution Signed-off-by: Jorge Albericio --- docs/developer/shared_prefix.md | 92 + .../core/extensions/transformer_engine.py | 66 +- megatron/core/model_parallel_config.py | 6 + megatron/core/models/hybrid/hybrid_model.py | 706 +++++++- .../hybrid/sequence_relative_attention.py | 132 ++ megatron/core/models/hybrid/shared_prefix.py | 1303 ++++++++++++++ .../core/models/hybrid/shared_prefix_fused.py | 1572 +++++++++++++++++ .../models/hybrid/shared_prefix_layout.py | 311 ++++ megatron/core/ssm/mamba_branch_layout.py | 113 ++ megatron/core/ssm/mamba_context_parallel.py | 19 +- megatron/core/ssm/mamba_forest_replay.py | 112 ++ megatron/core/ssm/mamba_mixer.py | 29 +- megatron/core/ssm/mamba_ragged.py | 265 +++ megatron/core/ssm/mamba_ragged_scan.py | 498 ++++++ megatron/core/ssm/mamba_sequence_packing.py | 136 ++ .../tensor_parallel/ordered_reduce_scatter.py | 93 + megatron/core/tensor_parallel/random.py | 28 +- megatron/core/transformer/attention.py | 45 +- megatron/core/transformer/moe/moe_layer.py | 19 +- megatron/core/transformer/moe/moe_utils.py | 116 +- megatron/core/transformer/moe/router.py | 122 +- .../transformer/multi_token_prediction.py | 44 +- .../core/transformer/transformer_config.py | 25 + .../test_shared_prefix_port_contracts.py | 130 ++ 24 files changed, 5847 insertions(+), 135 deletions(-) create mode 100644 docs/developer/shared_prefix.md create mode 100644 megatron/core/models/hybrid/sequence_relative_attention.py create mode 100644 megatron/core/models/hybrid/shared_prefix.py create mode 100644 megatron/core/models/hybrid/shared_prefix_fused.py create mode 100644 megatron/core/models/hybrid/shared_prefix_layout.py create mode 100644 megatron/core/ssm/mamba_branch_layout.py create mode 100644 megatron/core/ssm/mamba_forest_replay.py create mode 100644 megatron/core/ssm/mamba_ragged.py create mode 100644 megatron/core/ssm/mamba_ragged_scan.py create mode 100644 megatron/core/ssm/mamba_sequence_packing.py create mode 100644 megatron/core/tensor_parallel/ordered_reduce_scatter.py create mode 100644 tests/unit_tests/models/hybrid/test_shared_prefix_port_contracts.py diff --git a/docs/developer/shared_prefix.md b/docs/developer/shared_prefix.md new file mode 100644 index 00000000000..302ebc493df --- /dev/null +++ b/docs/developer/shared_prefix.md @@ -0,0 +1,92 @@ +# Shared-prefix execution for hybrid models + +This experimental, opt-in path avoids executing an identical prompt separately for +several completions. A caller supplies `shared_prefix_layout` to `HybridModel.forward`. +Without that argument, the ordinary model path, input-ID handling, quantization +initialization, and default execution settings remain in place. + +## Layout and attention + +A star stores `[prompt, completion_1, ..., completion_G]` once. A forest packs +several independent stars into one forward. Every completion attends to its prompt +and its own causal history. It cannot attend to a sibling completion or another +star. RoPE positions restart at the original prompt length for each completion. +The layout also records physical padding, logical lengths, and CP token ownership. +Backward accumulates the completion contributions into shared prompt activations. + +TP sequence parallelism and CP zigzag ownership use the caller's process groups. +The fused attention path exchanges the required sequence/head shards and retains +an independent causal domain for each completion. This is an execution layout; +logical sample weights, loss masks, and group normalization remain the caller's +responsibility. Log-probability and training forwards must use the same assignment +when comparing their outputs at unchanged weights. + +## Mamba state forking + +The recurrent prefix is evaluated up to a scan-chunk boundary. Its SSM state is +forked into independent completion branches. The remaining prompt tail is replayed +with each branch, and convolution carries the required prompt halo. This preserves +the kernel's chunk alignment and boundary context while eliminating most duplicated +prompt work. Ragged branches track their own lengths and boundaries rather than +turning the longest completion into useful work for every branch. Backward combines +branch state, halo, and shared-prefix contributions. + +The implementation therefore does not promise that every prompt token executes +exactly once in every Mamba sub-operation. The shared aligned prefix, residual-tail +replay, and convolution halo are distinct parts of the contract. + +## MoE, recomputation, and MTP + +Shared prompt rows carry their logical multiplicity for expert-bias statistics. +Both boolean routing maps and upstream dense top-k expert-index maps are supported; +padding and invalid dense routes contribute zero. Ordinary routing counts follow +the upstream path when multiplicity metadata is absent. Hash MoE remains supported +by the ordinary path and is explicitly rejected for shared execution. + +The router may run fixed row blocks inside a scoped shared forward. Activation +recomputation restores that scope while suppressing tensor-observation callbacks +inside the restored context, so a logical forward is observed only once. Frozen +router parameters skip unused parameter gradients while preserving gradients into +hidden states. + +MTP uses dense branch inputs and its existing heads. It does not share the MTP +prefix. When multiple independently normalized groups share one forward, optional +`loss_group_lengths` preserves their token-count correction. Ordinary MTP still +accepts precomputed decoder embeddings and upstream CP layout preparation. + +## Supported scope and explicit guards + +The shared adapter currently targets complete PP1 hybrid models with fp16/bf16, +zero dropout, RoPE or no positional embedding, ordinary self-attention and supported +Mamba/MoE layers. TP greater than one requires sequence parallelism. Full activation +recomputation supports the uniform method. Unsupported combinations fail explicitly, +including hash routing, quantization recipes, fp8/fp4, wide residual streams, mHC, +attention-logit softcapping, external attention/padding masks, inference contexts, +fine-grained activation offloading, CUDA graphs, sliding-window attention, auxiliary +router losses, and expert-capacity token dropping. See the validation functions for +the complete runtime contract. + +`sequence_relative_kernels` and `deterministic_tp_reduce_scatter` are separate, +default-off experimental controls. The former changes the ordinary attention/Mamba +numerical backend and must not be treated as a requirement for enabling the shared +layout. Neither is qualified merely because the shared-prefix model runs. + +## Validation status + +This current-main port has syntax, static-interface, and isolated CPU contract +checks for layout accounting, branch copy gradients, frozen-router gradients, and +recompute observation scope. These checks do not execute CUDA, Triton, +FlashAttention, Transformer Engine, distributed CP/TP/EP, or the full model. + +A previous production revision completed training and checkpoint evaluations. +That is evidence for that revision and configuration, not GPU qualification of this +port. Within-implementation forward agreement, cross-implementation gradients, +and long-run evaluation quality are separate acceptance criteria. The controlled +backbone-gradient discrepancy remains open; no gradient-identity or general +numerical-equivalence claim is made here. + +Before promotion: run native distributed tests and kernel replay coverage, compare +same-weight dense/shared outputs and gradients against repeated dense and repeated +shared controls, exercise matched packing and loss masks, validate MTP and expert +statistics, and run a bounded training/evaluation qualification. New kernel files +also require registration and replay tests in the upstream determinism manifest. diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index c362d72f804..6bb11d64821 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -42,6 +42,9 @@ set_tensor_model_parallel_attributes, ) from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region +from megatron.core.tensor_parallel.ordered_reduce_scatter import ( + ordered_reduce_scatter_to_sequence_parallel_region, +) from megatron.core.tensor_parallel.random import ( get_cuda_rng_tracker, get_data_parallel_rng_tracker_name, @@ -2057,6 +2060,28 @@ def __init__( self._tp_group = tp_group gtp_remat_group = resolve_gtp_remat_group(pg_collection, is_expert) + self._ordered_tp_reduce_scatter = ( + config.deterministic_tp_reduce_scatter and not is_expert and tp_group.size() > 1 + ) + self._ordered_reduce_add_bias = ( + self._ordered_tp_reduce_scatter and bias and not skip_bias_add + ) + if self._ordered_tp_reduce_scatter: + if not config.sequence_parallel: + raise ValueError("deterministic TP reduce-scatter requires sequence parallelism") + if config.tp_comm_overlap or config.symmetric_ar_type is not None: + raise ValueError( + "deterministic TP reduce-scatter does not support TP communication overlap" + ) + if config.fp8 or config.fp4 or config.quant_recipe is not None: + raise ValueError( + "deterministic TP reduce-scatter requires unquantized linear layers" + ) + if config.gtp_weight_remat_size != 1: + raise ValueError( + "deterministic TP reduce-scatter does not support weight rematerialization" + ) + super().__init__( input_size=input_size, output_size=output_size, @@ -2068,7 +2093,7 @@ def __init__( else lambda w: None ), bias=bias, - skip_bias_add=skip_bias_add, + skip_bias_add=True if self._ordered_tp_reduce_scatter else skip_bias_add, skip_weight_param_allocation=False, # We don't currently use this for row parallel layers # pylint: disable=line-too-long is_expert=is_expert, @@ -2079,6 +2104,11 @@ def __init__( gtp_remat_group=gtp_remat_group, gtp_replica_group=getattr(pg_collection, "expt_dp" if is_expert else "dp_cp", None), ) + if self._ordered_tp_reduce_scatter: + # Initialize as row-parallel to preserve weight sharding, RNG, and + # checkpoint attributes. Disable TE communication once, before any + # forward; our autograd reduction owns the matching backward gather. + self.parallel_mode = None if config.use_cpu_initialization: world_size = get_pg_size(tp_group) rank = get_pg_rank(tp_group) @@ -2113,6 +2143,18 @@ def __init__( ) _set_expert_parameter_attributes(self, "row", use_expert_pgs) + def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor | None]: + """Apply the local GEMM, then the configured sequence-parallel reduction.""" + output, bias = super().forward(x) + if self._ordered_tp_reduce_scatter: + output = ordered_reduce_scatter_to_sequence_parallel_region( + output, group=self._tp_group + ) + if self._ordered_reduce_add_bias: + output = output + bias + bias = None + return output, bias + def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): """Sharding along axis 1, bias not sharded""" state_dict = self.state_dict(prefix="", keep_vars=True) @@ -2443,6 +2485,28 @@ def _forward( bf16_backward: Optional[bool] = None, ) -> torch.Tensor: """Run TE attention after the runtime CP binding has been installed.""" + if self.config.sequence_relative_kernels: + if bf16_backward: + raise NotImplementedError( + "sequence-relative attention does not support a BF16-backward override" + ) + # This optional backend imports FlashAttention only when explicitly enabled. + from megatron.core.models.hybrid.sequence_relative_attention import ( + sequence_relative_attention_forward, + ) + + return sequence_relative_attention_forward( + self, + query, + key, + value, + attention_mask, + attn_mask_type, + attention_bias, + packed_seq_params, + num_splits, + ) + if packed_seq_params is not None: self.kept_packed_seq_params.discard("cp_group") self.kept_packed_seq_params.discard("local_cp_size") diff --git a/megatron/core/model_parallel_config.py b/megatron/core/model_parallel_config.py index 72e5000d44c..23e932d438c 100644 --- a/megatron/core/model_parallel_config.py +++ b/megatron/core/model_parallel_config.py @@ -97,6 +97,12 @@ class ModelParallelConfig: (https://arxiv.org/abs/2205.05198) for more details. """ + deterministic_tp_reduce_scatter: bool = False + """Use fixed-order FP32 accumulation for non-expert TE row-parallel outputs. + Requires sequence parallelism, unquantized GEMMs, and no TP communication + overlap or weight rematerialization. Changes rounding versus NCCL SUM. + """ + context_parallel_size: int = 1 """Splits network input along sequence dimension across GPU ranks.""" diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index d846d5aaa2a..cf0ad08bb53 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -1,6 +1,7 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging +from collections.abc import Iterator from contextlib import nullcontext from typing import Literal, Optional @@ -17,6 +18,14 @@ from megatron.core.models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.models.hybrid.layers import utils as layer_utils +from megatron.core.models.hybrid.shared_prefix import ( + SHARED_PREFIX_MTP_DENSE_HEADS_CAPABILITY, + SHARED_PREFIX_TRAINING_CAPABILITIES, + SharedPrefixForestLayout, + SharedPrefixLayout, + _validate_shared_prefix_physical_length, + forward_hybrid_stack_shared_prefix, +) from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( FineGrainedActivationOffloadingInterface as off_interface, @@ -26,7 +35,7 @@ from megatron.core.tensor_observation import observe_tensor from megatron.core.tensor_parallel import gather_from_sequence_parallel_region from megatron.core.transformer import TransformerConfig -from megatron.core.transformer.enums import InferenceCudaGraphScope, ModelType +from megatron.core.transformer.enums import AttnBackend, InferenceCudaGraphScope, ModelType from megatron.core.transformer.module import GraphableMegatronModule from megatron.core.transformer.moe.paged_stash import paged_stash_init_chunk_handler from megatron.core.transformer.multi_token_prediction import ( @@ -101,6 +110,262 @@ def _validate_hash_moe_pipeline_placement( ) +def _canonicalize_shared_prefix_cp_sequence( + local_tensor: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, + physical_len: int, + cp_group: torch.distributed.ProcessGroup, + *, + reduce_scatter_grad: bool, +) -> Tensor: + """Gather one zigzag CP shard and restore canonical global star order.""" + cp_size = cp_group.size() + if local_tensor.shape[0] * cp_size != physical_len: + raise ValueError( + "shared-prefix CP-local tensor length does not match the physical star: " + f"{local_tensor.shape[0]} * {cp_size} != {physical_len}" + ) + if cp_size == 1: + return local_tensor + + rank_order_tensor = gather_from_sequence_parallel_region( + local_tensor, tensor_parallel_output_grad=reduce_scatter_grad, group=cp_group + ) + rank_order_indices = torch.cat( + [ + layout.cp_local_indices(physical_len, cp_size, rank, local_tensor.device) + for rank in range(cp_size) + ] + ) + inverse_order = torch.empty_like(rank_order_indices) + inverse_order[rank_order_indices] = torch.arange( + physical_len, device=local_tensor.device, dtype=torch.long + ) + return rank_order_tensor.index_select(0, inverse_order) + + +def _validate_shared_prefix_mtp_star( + global_hidden_states: Tensor, + global_input_ids: Tensor, + global_loss_mask: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, + *, + cp_size: int, + cp_rank: int, +) -> None: + """Validate one canonical star before its dense MTP branches are reconstructed. + + The loss-mask contract (no loss on the prompt, on ordinary per-branch padding, + or on topology-only padding) is checked with one device-side reduction and a + single host sync. The offending region class is re-derived, with extra + syncs, only on the error path so the messages stay specific. + """ + physical_len = global_hidden_states.shape[0] + if global_hidden_states.ndim != 3 or global_hidden_states.shape[1] != 1: + raise ValueError("shared-prefix MTP hidden states must have shape [tokens, 1, hidden]") + if global_input_ids.shape != (1, physical_len): + raise ValueError("shared-prefix MTP input IDs must have shape [1, physical_tokens]") + if global_loss_mask.shape != (1, physical_len): + raise ValueError("shared-prefix MTP loss mask must have shape [1, physical_tokens]") + if physical_len < layout.total_len: + raise ValueError( + f"shared-prefix MTP physical length {physical_len} is shorter than layout " + f"length {layout.total_len}" + ) + if any( + root.logical_completion_lens is None or root.padding_multiple is None + for _, root in layout.iter_roots() + ): + raise NotImplementedError( + "shared-prefix MTP requires explicit logical completion lengths and physical padding" + ) + if cp_size < 1 or not 0 <= cp_rank < cp_size: + raise ValueError("shared-prefix MTP received an invalid CP size/rank") + + # Every region below is defined by the layout's Python ints, so the mask is + # built with plain kernels and the whole contract costs one ``.item()``. + device = global_loss_mask.device + must_be_zero = torch.ones(physical_len, dtype=torch.bool, device=device) + for offset, root in layout.iter_roots(): + for branch, logical_len in zip( + root.completion_slices(), root.logical_completion_lens, strict=True + ): + must_be_zero[offset + branch.start : offset + branch.start + logical_len] = False + if ((global_loss_mask[0] != 0) & must_be_zero).any().item(): + for offset, root in layout.iter_roots(): + if torch.count_nonzero(global_loss_mask[:, offset : offset + root.prefix_len]).item(): + raise ValueError("shared-prefix MTP loss mask must exclude every prompt token") + if torch.count_nonzero(global_loss_mask[:, layout.total_len :]).item(): + raise ValueError("shared-prefix MTP loss mask must exclude topology-only padding") + raise ValueError("shared-prefix MTP loss mask must exclude ordinary per-sequence padding") + + +def _shared_prefix_mtp_branch_indices( + layout: SharedPrefixLayout | SharedPrefixForestLayout, + device: torch.device | str, + *, + cp_size: int = 1, + cp_rank: int = 0, +) -> tuple[tuple[Tensor, ...], tuple[Tensor, ...]]: + """Return per-branch ``(star_indices, dense_positions)`` in this rank's CP-local order. + + Branch ``i`` is the conventional dense sequence ``[prompt, completion_i]``, + including that completion's ordinary per-sequence padding. ``star_indices[i]`` + selects those tokens from the canonical global star and ``dense_positions[i]`` + is each selected token's ``arange`` position inside its own dense branch. With + CP enabled both are composed with the branch's standard two-chunk zigzag + ownership, so the full dense branch is never allocated. + """ + star_indices = [] + dense_positions = [] + for indices in layout.dense_branch_indices(device): + if cp_size > 1: + positions = SharedPrefixLayout.cp_local_indices( + indices.numel(), cp_size, cp_rank, device + ) + indices = indices.index_select(0, positions) + else: + positions = torch.arange(indices.numel(), device=device, dtype=torch.long) + star_indices.append(indices) + dense_positions.append(positions) + return tuple(star_indices), tuple(dense_positions) + + +def _iter_shared_prefix_mtp_branches( + global_hidden_states: Tensor, + global_input_ids: Tensor, + global_loss_mask: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, + *, + cp_size: int = 1, + cp_rank: int = 0, +) -> Iterator[tuple[Tensor, Tensor, Tensor]]: + """Yield independent dense MTP branches from one canonical star. + + The main Hybrid stack can share its prompt, but MTP shifts future tokens and + therefore needs a conventional sequence boundary per completion. This is the + per-branch reference form of ``_pack_shared_prefix_mtp_branches``: the + production forward packs every branch into one sequence, while tests use + this iterator to check that packing against branch-by-branch reconstruction. + When CP is enabled, global-star indices are composed directly with that + branch's zigzag CP ownership so the full dense branch is never allocated. + """ + _validate_shared_prefix_mtp_star( + global_hidden_states, + global_input_ids, + global_loss_mask, + layout, + cp_size=cp_size, + cp_rank=cp_rank, + ) + star_indices, _ = _shared_prefix_mtp_branch_indices( + layout, global_hidden_states.device, cp_size=cp_size, cp_rank=cp_rank + ) + for indices in star_indices: + yield ( + global_hidden_states.index_select(0, indices), + global_input_ids.index_select(1, indices), + global_loss_mask.index_select(1, indices), + ) + + +def _pack_shared_prefix_mtp_branches( + global_hidden_states: Tensor, + global_input_ids: Tensor, + global_loss_mask: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, + *, + cp_size: int = 1, + cp_rank: int = 0, +) -> tuple[Tensor, Tensor, Tensor, Tensor]: + """Pack every dense MTP branch of one canonical star into a single THD sequence. + + Returns ``(hidden_states, input_ids, loss_mask, position_ids)`` for the + branch-major concatenation ``[prompt + completion_1 | ... | prompt + + completion_G]`` in this rank's CP-local order, gathered with exactly one + ``index_select`` per tensor. The result equals the concatenation of + ``_iter_shared_prefix_mtp_branches`` for the same CP rank; ``position_ids`` + restart at zero for every branch, exactly as a conventional dense batch. + """ + _validate_shared_prefix_mtp_star( + global_hidden_states, + global_input_ids, + global_loss_mask, + layout, + cp_size=cp_size, + cp_rank=cp_rank, + ) + star_indices, dense_positions = _shared_prefix_mtp_branch_indices( + layout, global_hidden_states.device, cp_size=cp_size, cp_rank=cp_rank + ) + packed_indices = torch.cat(star_indices) + return ( + global_hidden_states.index_select(0, packed_indices), + global_input_ids.index_select(1, packed_indices), + global_loss_mask.index_select(1, packed_indices), + torch.cat(dense_positions).unsqueeze(0), + ) + + +def _reconstruct_shared_prefix_mtp_branches( + global_hidden_states: Tensor, + global_input_ids: Tensor, + global_loss_mask: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, +) -> tuple[tuple[Tensor, Tensor, Tensor], ...]: + """Materialize the iterator for focused reconstruction/parity tests.""" + return tuple( + _iter_shared_prefix_mtp_branches( + global_hidden_states, global_input_ids, global_loss_mask, layout + ) + ) + + +def _validate_shared_prefix_mtp_pattern(mtp_pattern: Optional[str]) -> None: + """Keep the promoted predictor scope narrower than the Hybrid backbone scope.""" + if mtp_pattern is not None and 'M' in mtp_pattern: + raise NotImplementedError( + "shared-prefix MTP supports a Mamba Hybrid backbone but does not yet " + "support a Mamba layer inside the MTP predictor; use an attention/MLP " + "or attention/MoE MTP pattern such as '*-' or '*E'" + ) + + +def _validate_shared_prefix_mtp_attention_backend(attention_backend: AttnBackend) -> None: + """Reject the one attention backend that can never run the packed MTP branches. + + Shared-prefix MTP packs every dense branch into a single THD sequence. mcore's + local ``DotProductAttention`` asserts ``packed_seq_params is None``, so it cannot + serve that pack under any version. Transformer Engine backends are left to TE's + own per-release backend selection, which raises when the installed release has + no THD-capable kernel for the requested backend. + """ + if attention_backend == AttnBackend.local: + raise NotImplementedError( + "shared-prefix MTP packs every dense branch into one THD sequence, which " + "mcore's local DotProductAttention does not support; use a Transformer Engine " + "attention backend (flash/fused/unfused/auto)" + ) + + +def _hybrid_mtp_is_enabled( + configured_num_layers: Optional[int], mtp_pattern: Optional[str], mtp_pattern_depths: int +) -> bool: + """Return whether this Hybrid runtime should construct and execute MTP. + + Native Nemotron-H configs can retain MTP pattern metadata while a training + recipe explicitly overrides ``mtp_num_layers=0``. The pattern still + describes the checkpoint architecture, but zero must disable the runtime + MTP block and its post-processing path. + """ + return bool( + configured_num_layers is not None + and configured_num_layers > 0 + and mtp_pattern is not None + and mtp_pattern_depths > 0 + ) + + class HybridModel(LanguageModule, GraphableMegatronModule): """Hybrid language model. @@ -298,8 +563,9 @@ def __init__( # Determine if MTP is needed (based on pattern parsing) self.mtp_process = ( - self.mtp_pattern is not None - and self.mtp_num_depths > 0 + _hybrid_mtp_is_enabled( + self.config.mtp_num_layers, self.mtp_pattern, self.mtp_num_depths + ) # The following forces MTP to be on the final pipeline stage. It might be more optimal # to split the hybrid layer pattern into pipeline stages before parsing the pattern for # the current pipeline stage. This could also enable MTP standalone (MTP in a pipeline @@ -532,6 +798,196 @@ def create_mcore_cudagraph_manager(self, config): self.cudagraph_manager = CudaGraphManager(config) + def _forward_shared_prefix_mtp( + self, + *, + hidden_states: Tensor, + input_ids: Tensor, + loss_mask: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, + output_weight: Optional[Tensor], + runtime_gather_output: Optional[bool], + ) -> Tensor: + """Run the MTP heads once over all dense branches, preserving the backbone graph. + + The star is expanded into the branch-major dense sequence + ``[prompt + completion_1 | ... | prompt + completion_G]`` and the MTP block runs + a single time on it as a THD packed batch. Attention, RoPE and the MTP token + shifts all respect the ``cu_seqlens`` branch boundaries, so this equals running + the block once per branch while issuing one set of kernels and collectives. + """ + if ( + isinstance(layout, SharedPrefixForestLayout) + and layout.mtp_loss_group_root_counts + and not self.config.calculate_per_token_loss + ): + raise NotImplementedError( + "explicit MTP loss groups require calculate_per_token_loss=True" + ) + tp_group = self.pg_collection.tp + cp_group = self.pg_collection.cp + tp_size = tp_group.size() + cp_size = cp_group.size() + physical_len = input_ids.shape[1] * cp_size + + # The shared backbone is SP-sharded over TP. Gather without reducing + # gradients because the packed branches are scattered over TP again + # below, so each sequence position has one downstream owner. + # Detach before gathering and expansion when MTP cannot update the backbone. + # The loss anchor otherwise retains zero-gradient backward communication. + cp_local_hidden = hidden_states.detach() if self.config.mtp_detach_heads else hidden_states + if tp_size > 1: + cp_local_hidden = gather_from_sequence_parallel_region( + cp_local_hidden, tensor_parallel_output_grad=False, group=tp_group + ) + + global_hidden = _canonicalize_shared_prefix_cp_sequence( + cp_local_hidden, layout, physical_len, cp_group, reduce_scatter_grad=True + ) + global_input_ids = _canonicalize_shared_prefix_cp_sequence( + input_ids.transpose(0, 1).contiguous(), + layout, + physical_len, + cp_group, + reduce_scatter_grad=False, + ).transpose(0, 1) + global_loss_mask = _canonicalize_shared_prefix_cp_sequence( + loss_mask.transpose(0, 1).contiguous(), + layout, + physical_len, + cp_group, + reduce_scatter_grad=False, + ).transpose(0, 1) + global_branch_lengths = layout.dense_branch_lengths + if any(branch_len % cp_size for branch_len in global_branch_lengths): + raise ValueError("shared-prefix MTP branch length must be divisible by CP size") + combined_cp_length = sum(global_branch_lengths) // cp_size + if combined_cp_length % tp_size: + raise ValueError( + "shared-prefix MTP combined CP-local sequence must be divisible by TP size" + ) + packed_sp_length = combined_cp_length // tp_size + required_alignment = tp_size if cp_size == 1 else 2 * cp_size * tp_size + for branch_len in global_branch_lengths: + if branch_len % required_alignment: + raise ValueError( + "shared-prefix MTP physical branch length must be divisible by " + f"the CP/TP sequence quantum {required_alignment}, got {branch_len}" + ) + + # Build the dense branch-major sequence [P+C_1 | ... | P+C_G] in this + # rank's CP-local zigzag order with one index_select per tensor. Under + # THD every branch is its own attention sequence and its own roll_tensor + # segment, so one MTP call over the pack equals G independent dense calls. + packed_hidden, packed_input_ids, packed_loss_mask, packed_position_ids = ( + _pack_shared_prefix_mtp_branches( + global_hidden, + global_input_ids, + global_loss_mask, + layout, + cp_size=cp_size, + cp_rank=cp_group.rank(), + ) + ) + del cp_local_hidden, global_hidden, global_input_ids, global_loss_mask + if packed_hidden.shape[0] != combined_cp_length: + raise RuntimeError( + "shared-prefix MTP packed CP-local sequence has an invalid length: " + f"{packed_hidden.shape[0]} != {combined_cp_length}" + ) + if tp_size > 1: + # Scatter the whole pack once. This rank's SP shard is its contiguous + # slice of the branch-major sequence, which is the exact per-depth + # layout process_mtp_loss consumes below, so no TP gathers are needed. + packed_hidden = tensor_parallel.scatter_to_sequence_parallel_region( + packed_hidden, group=tp_group + ) + + # Global (pre-CP) cumulative branch lengths, as in the packed non-shared + # path: TE attention, THD RoPE, and roll_tensor divide by CP internally and + # map each branch onto this rank's two zigzag chunks. cp_group/local_cp_size + # stay unset exactly like trainer-built PackedSeqParams, so TE keeps the CP + # group it was constructed with instead of re-binding it every microbatch. + cumulative_lengths = [0] + for branch_len in global_branch_lengths: + cumulative_lengths.append(cumulative_lengths[-1] + branch_len) + cu_seqlens = torch.tensor( + cumulative_lengths, device=packed_input_ids.device, dtype=torch.int32 + ) + packed_seq_params = PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + max_seqlen_q=max(global_branch_lengths), + max_seqlen_kv=max(global_branch_lengths), + ) + + # Mirror the packed non-shared forward: one table covering the longest + # branch, not CP-sliced (packed_seq=True). THD RoPE selects each branch's + # positions from cu_seqlens, taking this CP rank's front/back zigzag chunks + # of every branch, which is what RotaryEmbedding(branch_len) did per branch. + if self.position_embedding_type == 'rope': + rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( + None, None, None, self.config, packed_seq_params + ) + packed_rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len, packed_seq=True) + elif self.position_embedding_type == 'none': + packed_rotary_pos_emb = None + else: + raise NotImplementedError( + "shared-prefix MTP supports only RoPE or positionless Hybrid models" + ) + packed_mtp_hidden = self.mtp( + input_ids=packed_input_ids, + position_ids=packed_position_ids, + hidden_states=packed_hidden, + attention_mask=None, + inference_params=None, + rotary_pos_emb=packed_rotary_pos_emb, + packed_seq_params=packed_seq_params, + embedding=self.embedding, + ) + del packed_hidden + # MultiTokenPredictionBlock returns depth-major chunks [depth_0 | ... | + # depth_D], each the SP shard of the branch-major pack, which is the layout + # process_mtp_loss chunks by depth. + depth_count = 1 + self.config.mtp_num_layers + if packed_mtp_hidden.shape[0] != depth_count * packed_sp_length: + raise RuntimeError( + "shared-prefix MTP returned an invalid depth-major SP shard length: " + f"{packed_mtp_hidden.shape[0]} != {depth_count} * {packed_sp_length}" + ) + processed_mtp_hidden = process_mtp_loss( + hidden_states=packed_mtp_hidden, + labels=None, + loss_mask=packed_loss_mask, + output_layer=self.output_layer, + output_weight=output_weight, + runtime_gather_output=runtime_gather_output, + is_training=self.training, + compute_language_model_loss=self.compute_language_model_loss, + config=self.config, + cp_group=cp_group, + tp_group=self.tp_group, + packed_seq_params=packed_seq_params, + scale_logits_fn=self._scale_logits if self.config.use_mup else None, + input_ids=packed_input_ids, + metric_avg_group=( + getattr(self.pg_collection, "dp_cp_gtp_remat", None) or self.pg_collection.dp_cp + ), + loss_group_lengths=( + tuple(length // cp_size for length in layout.mtp_loss_group_lengths) + if isinstance(layout, SharedPrefixForestLayout) + else None + ), + ) + + # The external RL loss consumes star logits. A zero-valued attachment + # retains process_mtp_loss's MTPLossAutoScaler nodes without changing + # those logits; its backward hook supplies the auxiliary-loss gradient. + mtp_loss_anchor = processed_mtp_hidden.reshape(-1)[0] * 0.0 + return hidden_states + mtp_loss_anchor + def forward( self, input_ids: Tensor, @@ -549,6 +1005,7 @@ def forward( padding_mask: Optional[Tensor] = None, compute_mtp_loss: bool = True, cp_batch: ContextParallelBatch | None = None, + shared_prefix_layout: Optional[SharedPrefixLayout | SharedPrefixForestLayout] = None, ) -> Tensor: """Forward function of the Hybrid model. This function passes the input tensors through the embedding layer, and then the decoder and finally into the post @@ -563,6 +1020,15 @@ def forward( ``labels`` still determine whether the model returns loss or logits. Defaults to True. cp_batch: Input tensors and packed metadata keyed by CP layout. + ``shared_prefix_layout`` explicitly selects the star path. The global packed input is + ``[prefix, completion_1, ..., completion_G, optional_topology_padding]`` with batch size + one; CP1 receives it whole, while CP>1 receives standard two-chunk zigzag sequence shards. + With TP>1, input IDs remain CP-local and replicated over TP; sequence parallelism starts at + the embedding output. The layout owns the exact tree mask and prefix-continued RoPE + positions. The normal decoder path is unchanged when the argument is ``None``. With + ``labels=None`` and parallel output, logits remain + ``[1, physical_len/CP, padded_vocab/TP]`` in the input shard's zigzag token order; this + model gathers neither sequence across CP nor vocabulary across TP. """ # If decoder_input is provided (not None), then input_ids and position_ids are ignored. # Otherwise, apply embedding layer on input_ids and position_ids to get decoder_input. @@ -577,6 +1043,87 @@ def forward( in_inference_mode = InferenceMode.is_active() + if shared_prefix_layout is not None: + if cp_batch is not None: + raise ValueError("Shared prefix owns CP token layout; cp_batch must be None") + if in_inference_mode or inference_context is not None: + raise NotImplementedError("shared-prefix Hybrid forward is training-only") + if attention_mask is not None: + raise ValueError( + 'shared-prefix Hybrid forward owns its tree mask and requires ' + 'attention_mask=None' + ) + if labels is not None: + raise NotImplementedError( + "shared-prefix Hybrid forward requires external next-token loss computation" + ) + if packed_seq_params is not None: + raise ValueError( + "shared_prefix_layout and packed_seq_params are mutually exclusive" + ) + if padding_mask is not None: + raise ValueError("shared-prefix Hybrid forward does not accept padding_mask") + if not self.pre_process or not self.post_process or self.vp_stage is not None: + raise NotImplementedError( + "shared-prefix HybridModel forward currently requires a complete PP1 model" + ) + if ( + self.mtp_process + and self.training + and SHARED_PREFIX_MTP_DENSE_HEADS_CAPABILITY + not in SHARED_PREFIX_TRAINING_CAPABILITIES + ): + raise NotImplementedError( + "shared-prefix Hybrid MTP dense-head support is implemented but not " + "advertised for production; distributed forward/backward parity must " + f"promote capability {SHARED_PREFIX_MTP_DENSE_HEADS_CAPABILITY!r}" + ) + if self.mtp_process and self.training: + _validate_shared_prefix_mtp_pattern(self.mtp_pattern) + _validate_shared_prefix_mtp_attention_backend(self.config.attention_backend) + if self.mtp_process and self.training and loss_mask is None: + raise ValueError("shared-prefix Hybrid MTP requires an explicit loss mask") + if self.position_embedding_type not in ('rope', 'none'): + raise NotImplementedError( + "shared-prefix Hybrid forward supports only RoPE or positionless attention" + ) + if self.config.multi_latent_attention: + raise NotImplementedError( + "shared-prefix Hybrid forward does not support multi-latent attention" + ) + tp_size = self.pg_collection.tp.size() + cp_size = self.pg_collection.cp.size() + if tp_size > 1 and not self.config.sequence_parallel: + raise NotImplementedError( + "shared-prefix HybridModel TP>1 requires sequence parallelism" + ) + if tp_size == 1 and self.config.sequence_parallel: + raise NotImplementedError( + "shared-prefix HybridModel sequence parallelism requires TP>1" + ) + if self.config.tensor_model_parallel_size != tp_size: + raise RuntimeError( + "shared-prefix HybridModel tensor-parallel config does not match its " + "process group" + ) + if tp_size > 1 and (not self.parallel_output or runtime_gather_output): + raise NotImplementedError( + "shared-prefix HybridModel TP/SP requires TP-sharded parallel output logits" + ) + if decoder_input is None: + if input_ids is None or input_ids.ndim != 2 or input_ids.shape[0] != 1: + raise ValueError( + "shared-prefix Hybrid input_ids must have shape [1, physical_len/CP]" + ) + physical_len = input_ids.shape[1] * cp_size + _validate_shared_prefix_physical_length( + shared_prefix_layout, + physical_len, + tp_size=tp_size, + cp_size=cp_size, + sequence_parallel=self.config.sequence_parallel, + ) + if in_inference_mode: assert runtime_gather_output, "Inference must always gather TP logits" @@ -648,7 +1195,32 @@ def forward( ) rotary_pos_emb = None - if self.position_embedding_type == 'rope' and not self.config.multi_latent_attention: + if shared_prefix_layout is not None and self.position_embedding_type == 'rope': + if decoder_input is None: + raise RuntimeError("shared-prefix Hybrid embedding did not produce decoder input") + cp_group = self.pg_collection.cp + cp_size = cp_group.size() + tp_size = self.pg_collection.tp.size() + sequence_shards = tp_size if self.config.sequence_parallel else 1 + physical_len = decoder_input.shape[0] * cp_size * sequence_shards + if decoder_input.ndim != 3 or decoder_input.shape[1] != 1: + raise ValueError( + "shared-prefix Hybrid decoder input must have shape " + "[physical_len/(TP*CP), 1, hidden] when sequence parallelism is enabled" + ) + rotary_table = self.rotary_pos_emb.get_emb( + max(shared_prefix_layout.dense_branch_lengths) + ) + global_position_ids = shared_prefix_layout.padded_position_ids( + physical_len, rotary_table.device + ) + if cp_size > 1: + local_indices = shared_prefix_layout.cp_local_indices( + physical_len, cp_size, cp_group.rank(), rotary_table.device + ) + global_position_ids = global_position_ids.index_select(0, local_indices) + rotary_pos_emb = rotary_table.index_select(0, global_position_ids) + elif self.position_embedding_type == 'rope' and not self.config.multi_latent_attention: rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( inference_context, self.decoder, decoder_input, self.config, packed_seq_params ) @@ -694,22 +1266,32 @@ def forward( else nullcontext() ) with backbone_context: - decoder_output = self.decoder( - hidden_states=decoder_input, - attention_mask=attention_mask, - inference_context=inference_context, - rotary_pos_emb=rotary_pos_emb, - packed_seq_params=packed_seq_params, - padding_mask=padding_mask, - packed_seq_params_by_layout=packed_seq_params_by_layout, - cp_layout_plan=cp_layout_plan, - input_ids=hash_input_ids, - ) - if isinstance(decoder_output, tuple): - hidden_states, mhc_multistream = decoder_output - else: - hidden_states = decoder_output - mhc_multistream = None + if shared_prefix_layout is not None: + hidden_states = forward_hybrid_stack_shared_prefix( + self.decoder, + decoder_input, + shared_prefix_layout, + rotary_pos_emb=rotary_pos_emb, + position_embedding_type=self.position_embedding_type, + ) + mhc_multistream = None + else: + decoder_output = self.decoder( + hidden_states=decoder_input, + attention_mask=attention_mask, + inference_context=inference_context, + rotary_pos_emb=rotary_pos_emb, + packed_seq_params=packed_seq_params, + padding_mask=padding_mask, + packed_seq_params_by_layout=packed_seq_params_by_layout, + cp_layout_plan=cp_layout_plan, + input_ids=hash_input_ids, + ) + if isinstance(decoder_output, tuple): + hidden_states, mhc_multistream = decoder_output + else: + hidden_states = decoder_output + mhc_multistream = None output_weight = None if self.share_embeddings_and_output_weights: @@ -726,48 +1308,64 @@ def forward( ) mtp_forward_ran = ( - self.mtp_process and not (in_inference_mode or is_spec_decode) and compute_mtp_loss + self.mtp_process + and not (in_inference_mode or is_spec_decode) + and compute_mtp_loss + and (shared_prefix_layout is None or self.training) ) mtp_hidden_states = hidden_states mtp_inputs = None if mtp_forward_ran: - mtp_inputs = self.mtp.prepare_cp_layout( - input_ids=input_ids, - position_ids=position_ids, - hidden_states=hidden_states, - decoder_input=decoder_input if use_precomputed_mtp_embeddings else None, - mhc_multistream=mhc_multistream, - labels=labels, - loss_mask=loss_mask, - mtp_input_mask=mtp_input_mask, - packed_seq_params=packed_seq_params, - cp_batch=cp_batch, - ) - if mtp_inputs.decoder_input is None: - assert mtp_inputs.input_ids is not None and mtp_inputs.position_ids is not None, ( - "MTP requires both input_ids and position_ids when precomputed " - "decoder_input embeddings are not provided." + if shared_prefix_layout is not None: + hidden_states = self._forward_shared_prefix_mtp( + hidden_states=hidden_states, + input_ids=input_ids, + loss_mask=loss_mask, + layout=shared_prefix_layout, + output_weight=output_weight, + runtime_gather_output=runtime_gather_output, + ) + mtp_hidden_states = hidden_states + else: + mtp_inputs = self.mtp.prepare_cp_layout( + input_ids=input_ids, + position_ids=position_ids, + hidden_states=hidden_states, + decoder_input=decoder_input if use_precomputed_mtp_embeddings else None, + mhc_multistream=mhc_multistream, + labels=labels, + loss_mask=loss_mask, + mtp_input_mask=mtp_input_mask, + packed_seq_params=packed_seq_params, + cp_batch=cp_batch, + ) + if mtp_inputs.decoder_input is None: + assert ( + mtp_inputs.input_ids is not None and mtp_inputs.position_ids is not None + ), ( + "MTP requires both input_ids and position_ids when precomputed " + "decoder_input embeddings are not provided." + ) + mtp_hidden_states = self.mtp( + input_ids=mtp_inputs.input_ids, + position_ids=mtp_inputs.position_ids, + hidden_states=mtp_inputs.hidden_states, + mhc_multistream=mtp_inputs.mhc_multistream, + attention_mask=attention_mask, + inference_params=inference_params, + rotary_pos_emb=rotary_pos_emb, + packed_seq_params=mtp_inputs.packed_seq_params, + embedding=self.embedding, + decoder_input=mtp_inputs.decoder_input, + mtp_input_mask=mtp_inputs.mtp_input_mask, + packed_seq_params_by_layout=packed_seq_params_by_layout, + cp_layout_plan=cp_layout_plan, ) - mtp_hidden_states = self.mtp( - input_ids=mtp_inputs.input_ids, - position_ids=mtp_inputs.position_ids, - hidden_states=mtp_inputs.hidden_states, - mhc_multistream=mtp_inputs.mhc_multistream, - attention_mask=attention_mask, - inference_params=inference_params, - rotary_pos_emb=rotary_pos_emb, - packed_seq_params=mtp_inputs.packed_seq_params, - embedding=self.embedding, - decoder_input=mtp_inputs.decoder_input, - mtp_input_mask=mtp_inputs.mtp_input_mask, - packed_seq_params_by_layout=packed_seq_params_by_layout, - cp_layout_plan=cp_layout_plan, - ) if not self.post_process: return mtp_hidden_states if mtp_forward_ran else hidden_states - if self.config.mtp_num_layers is not None and self.mtp_process: + if self.mtp_process: assert self.config.mtp_num_layers > 0 if is_spec_decode: assert inference_context is not None @@ -783,7 +1381,7 @@ def forward( # Non-block scope: direct assignment; the controller will set # this back to None after reading to allow GC. inference_context.mtp_decoder_hidden_states = hidden_states - elif mtp_forward_ran: + elif mtp_forward_ran and shared_prefix_layout is None: assert mtp_inputs is not None # For RL (labels is None), process_mtp_loss derives labels from # input_ids to match the SFT label format. diff --git a/megatron/core/models/hybrid/sequence_relative_attention.py b/megatron/core/models/hybrid/sequence_relative_attention.py new file mode 100644 index 00000000000..7c38489de4c --- /dev/null +++ b/megatron/core/models/hybrid/sequence_relative_attention.py @@ -0,0 +1,132 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Experimental dense CP adapter for the same unsplit FlashAttention arithmetic. + +This defines a new numerical baseline. It intentionally does not reproduce the +old TE ring's length-dependent partial-output rounding. +""" + +import itertools +from functools import lru_cache +from typing import TYPE_CHECKING + +import torch +from flash_attn import flash_attn_varlen_func + +from megatron.core.models.hybrid import shared_prefix_fused as attention +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.transformer.enums import AttnMaskType + +if TYPE_CHECKING: + from megatron.core.extensions.transformer_engine import TEDotProductAttention + + +@lru_cache(maxsize=128) +def rank_major_mapping( + boundaries: tuple[int, ...], cp_size: int, device: torch.device +) -> tuple[torch.Tensor, torch.Tensor]: + """Map CP rank-major token ownership to canonical sequence order and back.""" + indices = [] + for rank in range(cp_size): + for start, end in itertools.pairwise(boundaries): + assert (end - start) % (2 * cp_size) == 0 + segment = (end - start) // (2 * cp_size) + indices.extend( + ( + torch.arange( + start + rank * segment, start + (rank + 1) * segment, device=device + ), + torch.arange(end - (rank + 1) * segment, end - rank * segment, device=device), + ) + ) + forward = torch.cat(indices) + assert forward.numel() == boundaries[-1] + return forward, forward.argsort() + + +def dense_attention_cp( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + *, + cu_seqlens: torch.Tensor, + cp_group: torch.distributed.ProcessGroup, + scale: float | None = None, +) -> torch.Tensor: + """Exchange sequence shards for head shards, then attend each full sequence.""" + assert query.ndim == key.ndim == value.ndim == 3 + cp_size = cp_group.size() + boundaries = tuple(cu_seqlens.tolist()) + assert boundaries[0] == 0 and query.shape[0] * cp_size == boundaries[-1] + assert key.shape[0] == value.shape[0] == query.shape[0] + assert key.shape[1] == value.shape[1] + assert query.shape[2] == key.shape[2] == value.shape[2] + to_rank_major, to_canonical = rank_major_mapping(boundaries, cp_size, query.device) + slices = attention._cp_kv_head_slices_for_destinations(query.shape[1], key.shape[1], cp_size) + + def exchange(tensor, *, kv=False): + if kv: + tensor = torch.cat([tensor[:, part, :] for part in slices], dim=1) + assert tensor.shape[1] % cp_size == 0 + heads = tensor.shape[1] // cp_size + result = attention.all_to_all_sp2hp(tensor.reshape(tensor.shape[0], 1, -1), group=cp_group) + return result.reshape(boundaries[-1], heads, tensor.shape[-1])[to_canonical] + + q, k, v = exchange(query), exchange(key, kv=True), exchange(value, kv=True) + maximum = max(end - start for start, end in itertools.pairwise(boundaries)) + output = flash_attn_varlen_func( + q, + k, + v, + cu_seqlens, + cu_seqlens, + maximum, + maximum, + causal=True, + softmax_scale=scale, + deterministic=True, + ) + rank_major = output[to_rank_major].reshape(boundaries[-1], 1, -1) + result = attention.all_to_all_hp2sp(rank_major, group=cp_group) + return result.reshape(query.shape[0], query.shape[1] * value.shape[-1]) + + +def sequence_relative_attention_forward( + self: "TEDotProductAttention", + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attention_mask: torch.Tensor | None, + attn_mask_type: AttnMaskType, + attention_bias: torch.Tensor | None = None, + packed_seq_params: PackedSeqParams | None = None, + num_splits: int | None = None, +) -> torch.Tensor: + """Use the common sequence-relative backend for supported causal THD attention.""" + if packed_seq_params is None or packed_seq_params.qkv_format != "thd": + raise ValueError("sequence-relative attention requires packed THD input") + if attention_mask is not None or attention_bias is not None or num_splits is not None: + raise ValueError( + "sequence-relative attention does not support mask/bias/num_splits overrides" + ) + assert attn_mask_type.name in ("causal", "padding_causal") + assert self.config.attention_dropout == 0 and self.config.window_size is None + assert not self.config.qk_clip and not self.config.log_max_attention_logit + assert self.config.softmax_type == "vanilla" + assert packed_seq_params.local_cp_size is None + group = packed_seq_params.cp_group + if group is None: + group = self.cp_group + assert group is not None + # Shared-prefix MTP already aligns each dense branch to the CP/TP quantum + # and supplies cu_seqlens without the optional *_padded fields. + cu_q = packed_seq_params.cu_seqlens_q_padded + cu_k = packed_seq_params.cu_seqlens_kv_padded + if cu_q is None: + cu_q = packed_seq_params.cu_seqlens_q + if cu_k is None: + cu_k = packed_seq_params.cu_seqlens_kv + assert cu_q is not None and cu_k is not None and torch.equal(cu_q, cu_k) + return dense_attention_cp( + query, key, value, cu_seqlens=cu_q, cp_group=group, scale=self.config.softmax_scale + ) diff --git a/megatron/core/models/hybrid/shared_prefix.py b/megatron/core/models/hybrid/shared_prefix.py new file mode 100644 index 00000000000..7ff37013fc8 --- /dev/null +++ b/megatron/core/models/hybrid/shared_prefix.py @@ -0,0 +1,1303 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Explicit shared-prefix adapter for a HybridStack. + +This module deliberately does not alter :class:`HybridStack`'s normal ``forward`` path. Call +``forward_hybrid_stack_shared_prefix`` with one packed star ``[P, C_1, ..., C_G]`` to opt in. +Mamba layers scan the prefix once and fork differentiable convolution/SSM state into the +completion +branches, including CP>1 through sequence-to-head collectives. An exact uninterrupted-replay path +is retained as an explicit parity oracle and fallback. Attention uses the exact-backward fused +forest +kernel. +Unsupported topology or model features fail before executing a partial forward. + +The implementation is the narrow production slice of the shared-prefix work developed in +Megatron-RL/Megatron-LM commit ``0bf30804f`` plus the fused kernel port ``5b7173f7``. CP1 and CP>1 +are advertised as distinct capabilities so integrations can negotiate against the validated +topology. CP>1 uses the model's standard zigzag sequence shards and sequence-to-head all-to-alls. +Validated production capabilities cover TP1 and TP>1 with sequence parallelism, explicit physical +padding, MoE expert-bias accounting, and full uniform activation recomputation. Topology and +feature +tokens remain distinct so integrations can require the exact supported conjunction. +""" + +import os + +import torch +import torch.nn.functional as F +from einops import rearrange +from torch import Tensor + +from megatron.core import tensor_parallel +from megatron.core.models.hybrid.shared_prefix_layout import ( + SharedPrefixForestLayout, + SharedPrefixLayout, +) +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.ssm.mamba_branch_layout import merge_mamba_branches, pack_mamba_branches +from megatron.core.ssm.mamba_layer import MambaLayer +from megatron.core.ssm.mamba_mixer import ( + MAMBA_HAS_STATE_DTYPE, + MambaMixer, + causal_conv1d_fn, + mamba_chunk_scan_combined, +) +from megatron.core.transformer.attention import SelfAttention +from megatron.core.transformer.identity_op import IdentityOp +from megatron.core.transformer.moe.moe_utils import router_gating_token_blocks +from megatron.core.transformer.transformer_layer import TransformerLayer +from megatron.core.typed_torch import apply_module +from megatron.core.utils import make_viewless_tensor + + +def _mamba_state_dtype_kwargs(mixer: MambaMixer) -> dict: + if not MAMBA_HAS_STATE_DTYPE: + return {} + return {"state_dtype": mixer.mamba_training_ssm_states_dtype} + + +def _validate_mamba_fork(mixer: MambaMixer) -> None: + tp_size = mixer.pg_collection.tp.size() + if tp_size > 1 and not mixer.config.sequence_parallel: + raise NotImplementedError("shared-prefix Mamba TP>1 requires sequence parallelism") + if tp_size == 1 and mixer.config.sequence_parallel: + raise NotImplementedError("shared-prefix Mamba sequence parallelism requires TP>1") + if not mixer.rmsnorm: + raise NotImplementedError("shared-prefix Mamba state forking requires gated RMSNorm") + if causal_conv1d_fn is None or mamba_chunk_scan_combined is None: + raise RuntimeError("shared-prefix Mamba state forking requires causal-conv1d and mamba-ssm") + + +def _prefix_conv_context(xbc: Tensor, width: int) -> Tensor: + """Last ``width`` pre-convolution columns, including the causal left-zero padding.""" + if width == 0: + return xbc[:, :, :0] + padded = F.pad(xbc, (max(0, width - xbc.shape[-1]), 0)) + return padded[:, :, -width:].clone() + + +def _mamba_prefix_fork_boundary(mixer: MambaMixer, prefix_len: int) -> int: + """Return the last scan-chunk boundary at or before ``prefix_len``. + + Restarting ``mamba_chunk_scan_combined`` at an arbitrary token changes its + internal chunk partition. With BF16 training state, that extra state + boundary is numerically visible after many Hybrid layers. Fork only at an + ordinary Mamba chunk boundary and replay the (short) prompt tail inside + each completion branch so the state-fork and uninterrupted scans use the + same chunk partition. + """ + prefix_len = int(prefix_len) + chunk_size = getattr(mixer, "chunk_size", None) + if isinstance(chunk_size, bool) or not isinstance(chunk_size, int) or chunk_size < 1: + message = "shared-prefix Mamba requires a positive integer chunk_size" + raise ValueError(f"{message}, got {chunk_size!r}") + if prefix_len < 1: + raise ValueError("shared-prefix Mamba requires a non-empty prefix") + return prefix_len - prefix_len % chunk_size + + +def _scan_mamba_projected_segment( + mixer: MambaMixer, + projected: Tensor, + *, + conv_context: Tensor | None = None, + ssm_initial_state: Tensor | None = None, + capture_state: bool = False, + include_gate: bool = True, +) -> tuple[Tensor, Tensor | None, Tensor | None, Tensor | None]: + """Scan already CP-canonical, head-sharded Mamba projections for one segment batch.""" + cp = mixer.cp + if projected.ndim != 3: + raise ValueError( + "shared-prefix projected Mamba segments must have shape [sequence,batch,d]" + ) + num_groups = cp.ngroups_local_tpcp + num_heads = cp.nheads_local_tpcp + d_inner = cp.d_inner_local_tpcp + branch_count = projected.shape[1] + + projected = rearrange(projected, "l b d -> b l d").contiguous() + if include_gate: + z, projected = projected.split([d_inner, projected.shape[-1] - d_inner], dim=-1) + else: + z = None + xbc, dt = torch.split(projected, [d_inner + 2 * num_groups * mixer.d_state, num_heads], dim=-1) + A = -torch.exp(cp.get_A_log().float()) + + xbc = rearrange(xbc, "b l d -> b d l").contiguous() + next_conv_context = _prefix_conv_context(xbc, mixer.d_conv - 1) if capture_state else None + if conv_context is not None: + if ( + conv_context.ndim != 3 + or conv_context.shape[1] != xbc.shape[1] + or conv_context.shape[2] != mixer.d_conv - 1 + ): + raise ValueError("prefix convolution state is incompatible with the branch") + if conv_context.shape[0] not in (1, branch_count): + raise ValueError("prefix convolution state batch is incompatible with the branch") + repeated_context = conv_context.to(xbc.dtype).expand(branch_count, -1, -1) + conv_input = torch.cat([repeated_context, xbc], dim=-1) + conv_output = causal_conv1d_fn( + conv_input, + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + activation=mixer.activation, + )[:, :, repeated_context.shape[-1] :] + else: + conv_output = causal_conv1d_fn( + xbc, + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + activation=mixer.activation, + ) + xbc = rearrange(conv_output, "b d l -> b l d").contiguous() + + x, B, C = torch.split( + xbc, [d_inner, num_groups * mixer.d_state, num_groups * mixer.d_state], dim=-1 + ) + x = rearrange(x, "b l (h p) -> b l h p", p=mixer.headdim).contiguous() + B = rearrange(B, "b l (g n) -> b l g n", n=mixer.d_state).contiguous() + C = rearrange(C, "b l (g n) -> b l g n", n=mixer.d_state).contiguous() + + initial_states = ssm_initial_state + if initial_states is not None: + if initial_states.shape[0] not in (1, branch_count): + raise ValueError("prefix SSM state batch is incompatible with the branch") + initial_states = initial_states.expand(branch_count, *initial_states.shape[1:]).contiguous() + scan = mamba_chunk_scan_combined( + x, + dt.contiguous(), + A, + B, + C, + mixer.chunk_size, + D=( + rearrange(cp.get_D().float(), "(h p) -> h p", p=mixer.headdim) + if mixer.D_has_hdim + else cp.get_D() + ), + z=None, + dt_bias=cp.get_dt_bias().float(), + dt_softplus=True, + initial_states=initial_states, + return_final_states=capture_state, + **_mamba_state_dtype_kwargs(mixer), + ) + if capture_state: + y, final_state = scan + else: + y, final_state = scan, None + y = rearrange(y, "b l h p -> l b (h p)").contiguous() + if z is not None: + z = rearrange(z, "b l d -> l b d").contiguous() + return y, z, next_conv_context, final_state + + +def _fork_mamba_segment( + mixer: MambaMixer, + hidden_states: Tensor, + *, + conv_context: Tensor | None = None, + ssm_initial_state: Tensor | None = None, + capture_state: bool = False, +) -> tuple[Tensor, Tensor | None, Tensor | None, Tensor | None]: + """Scan one segment, optionally capturing a differentiable Mamba end state.""" + _validate_mamba_fork(mixer) + if mixer.pg_collection.tp.size() != 1: + raise RuntimeError( + "TP sequence-sharded Mamba segments require the collective parallel adapter path" + ) + if mixer.cp.cp_size != 1: + raise RuntimeError("CP-sharded Mamba segments require the collective CP adapter path") + if hidden_states.ndim != 3 or hidden_states.shape[1] != 1: + raise ValueError("shared-prefix Mamba segments must have shape [sequence, 1, hidden]") + + cp = mixer.cp + num_groups = cp.ngroups_local_tpcp + num_heads = cp.nheads_local_tpcp + d_inner = cp.d_inner_local_tpcp + + projected, _ = mixer.in_proj(hidden_states) + projected = cp.pre_conv_ssm(projected) + projected = rearrange(projected, "l b d -> b l d").contiguous() + z, xbc, dt = torch.split( + projected, [d_inner, d_inner + 2 * num_groups * mixer.d_state, num_heads], dim=-1 + ) + A = -torch.exp(cp.get_A_log().float()) + + xbc = rearrange(xbc, "b l d -> b d l").contiguous() + next_conv_context = _prefix_conv_context(xbc, mixer.d_conv - 1) if capture_state else None + if conv_context is not None: + if conv_context.shape[:2] != xbc.shape[:2]: + raise ValueError("prefix convolution state is incompatible with the branch") + conv_input = torch.cat([conv_context.to(xbc.dtype), xbc], dim=-1) + conv_output = causal_conv1d_fn( + conv_input, + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + activation=mixer.activation, + )[:, :, conv_context.shape[-1] :] + else: + conv_output = causal_conv1d_fn( + xbc, + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + activation=mixer.activation, + ) + xbc = rearrange(conv_output, "b d l -> b l d").contiguous() + + x, B, C = torch.split( + xbc, [d_inner, num_groups * mixer.d_state, num_groups * mixer.d_state], dim=-1 + ) + x = rearrange(x, "b l (h p) -> b l h p", p=mixer.headdim).contiguous() + B = rearrange(B, "b l (g n) -> b l g n", n=mixer.d_state).contiguous() + C = rearrange(C, "b l (g n) -> b l g n", n=mixer.d_state).contiguous() + z = rearrange(z, "b l (h p) -> b l h p", p=mixer.headdim).contiguous() + + scan = mamba_chunk_scan_combined( + x, + dt.contiguous(), + A, + B, + C, + mixer.chunk_size, + D=( + rearrange(cp.get_D().float(), "(h p) -> h p", p=mixer.headdim) + if mixer.D_has_hdim + else cp.get_D() + ), + z=None, + dt_bias=cp.get_dt_bias().float(), + dt_softplus=True, + initial_states=ssm_initial_state, + return_final_states=capture_state, + **_mamba_state_dtype_kwargs(mixer), + ) + if capture_state: + y, final_state = scan + else: + y, final_state = scan, None + + y = rearrange(y, "b l h p -> l b (h p)").contiguous() + y = cp.post_conv_ssm(y) + z = rearrange(z, "b l h p -> l b (h p)").contiguous() + z = cp.post_conv_ssm(z) + y = mixer.norm(y, z) + output, output_bias = mixer.out_proj(y) + return output, output_bias, next_conv_context, final_state + + +def _fork_mamba_branches( + mixer: MambaMixer, + branches: Tensor, + *, + conv_context: Tensor | None, + ssm_initial_state: Tensor | None, +) -> tuple[Tensor, Tensor | None]: + """Scan right-padded prompt-tail/completion branches from an aligned state.""" + _validate_mamba_fork(mixer) + if mixer.pg_collection.tp.size() != 1: + raise RuntimeError( + "TP sequence-sharded Mamba branches require the collective parallel adapter path" + ) + if mixer.cp.cp_size != 1: + raise RuntimeError("CP-sharded Mamba branches require the collective CP adapter path") + if branches.ndim != 3: + raise ValueError( + "shared-prefix Mamba branches must have shape [sequence, branches, hidden]" + ) + + cp = mixer.cp + num_groups = cp.ngroups_local_tpcp + num_heads = cp.nheads_local_tpcp + d_inner = cp.d_inner_local_tpcp + branch_count = branches.shape[1] + + projected, _ = mixer.in_proj(branches) + projected = cp.pre_conv_ssm(projected) + projected = rearrange(projected, "l b d -> b l d").contiguous() + z, xbc, dt = torch.split( + projected, [d_inner, d_inner + 2 * num_groups * mixer.d_state, num_heads], dim=-1 + ) + A = -torch.exp(cp.get_A_log().float()) + + xbc = rearrange(xbc, "b l d -> b d l").contiguous() + if conv_context is None: + conv_output = causal_conv1d_fn( + xbc, + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + activation=mixer.activation, + ) + else: + repeated_context = conv_context.to(xbc.dtype).expand(branch_count, -1, -1) + conv_input = torch.cat([repeated_context, xbc], dim=-1) + conv_output = causal_conv1d_fn( + conv_input, + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + activation=mixer.activation, + )[:, :, repeated_context.shape[-1] :] + xbc = rearrange(conv_output, "b d l -> b l d").contiguous() + + x, B, C = torch.split( + xbc, [d_inner, num_groups * mixer.d_state, num_groups * mixer.d_state], dim=-1 + ) + x = rearrange(x, "b l (h p) -> b l h p", p=mixer.headdim).contiguous() + B = rearrange(B, "b l (g n) -> b l g n", n=mixer.d_state).contiguous() + C = rearrange(C, "b l (g n) -> b l g n", n=mixer.d_state).contiguous() + z = rearrange(z, "b l (h p) -> b l h p", p=mixer.headdim).contiguous() + initial_states = ( + None + if ssm_initial_state is None + else ssm_initial_state.expand(branch_count, *ssm_initial_state.shape[1:]).contiguous() + ) + + y = mamba_chunk_scan_combined( + x, + dt.contiguous(), + A, + B, + C, + mixer.chunk_size, + D=( + rearrange(cp.get_D().float(), "(h p) -> h p", p=mixer.headdim) + if mixer.D_has_hdim + else cp.get_D() + ), + z=None, + dt_bias=cp.get_dt_bias().float(), + dt_softplus=True, + initial_states=initial_states, + return_final_states=False, + **_mamba_state_dtype_kwargs(mixer), + ) + y = rearrange(y, "b l h p -> l b (h p)").contiguous() + y = cp.post_conv_ssm(y) + z = rearrange(z, "b l h p -> l b (h p)").contiguous() + z = cp.post_conv_ssm(z) + y = mixer.norm(y, z) + return mixer.out_proj(y) + + +def _forward_mamba_layer_shared_prefix( + layer: MambaLayer, hidden_states: Tensor, layout: SharedPrefixLayout +) -> Tensor: + residual = hidden_states.float() if layer.config.fp32_residual_connection else hidden_states + normalized = apply_module(layer.norm)(hidden_states.to(dtype=layer.config.params_dtype)) + + fork_boundary = _mamba_prefix_fork_boundary(layer.mixer, layout.prefix_len) + replayed_prefix_len = layout.prefix_len - fork_boundary + prefix_output = normalized.new_empty((0, 1, normalized.shape[-1])) + output_bias = None + conv_context = None + final_state = None + if fork_boundary: + prefix_output, output_bias, conv_context, final_state = _fork_mamba_segment( + layer.mixer, normalized[:fork_boundary], capture_state=True + ) + physical_completion_lens = list(layout.completion_lens) + physical_completion_lens[-1] += hidden_states.shape[0] - layout.total_len + max_completion_len = replayed_prefix_len + max(physical_completion_lens) + branches = normalized.new_zeros( + max_completion_len, len(physical_completion_lens), normalized.shape[-1] + ) + start = layout.prefix_len + for branch_index, completion_len in enumerate(physical_completion_lens): + if replayed_prefix_len: + branches[:replayed_prefix_len, branch_index] = normalized[ + fork_boundary : layout.prefix_len, 0 + ] + branches[replayed_prefix_len : replayed_prefix_len + completion_len, branch_index] = ( + normalized[start : start + completion_len, 0] + ) + start += completion_len + if start != hidden_states.shape[0]: + raise RuntimeError("shared-prefix branch spans do not cover physical sequence") + branch_output, branch_bias = _fork_mamba_branches( + layer.mixer, branches, conv_context=conv_context, ssm_initial_state=final_state + ) + if output_bias is None: + output_bias = branch_bias + prefix_output = torch.cat([prefix_output, branch_output[:replayed_prefix_len, :1]], dim=0) + packed_output = torch.cat( + [prefix_output] + + [ + branch_output[replayed_prefix_len : replayed_prefix_len + length, index : index + 1] + for index, length in enumerate(physical_completion_lens) + ], + dim=0, + ) + + with layer.bias_dropout_add_exec_handler(): + return layer.mamba_bda(training=layer.training, fused=layer.config.bias_dropout_fusion)( + (packed_output, output_bias), residual, layer.hidden_dropout + ) + + +def _scan_mamba_shared_prefix_root( + mixer: MambaMixer, projected: Tensor, layout: SharedPrefixLayout, *, replay_prefix: bool +) -> Tensor: + """Scan recurrent fields for one root, retaining the established chunk boundary.""" + physical_len = projected.shape[0] + physical_completion_lens = list(layout.completion_lens) + # The topology validator owns physical alignment. Padding is placed after the final + # completion; + # scanning it as that branch's causal tail cannot affect any real-token output. + physical_completion_lens[-1] += physical_len - layout.total_len + fork_boundary = 0 if replay_prefix else _mamba_prefix_fork_boundary(mixer, layout.prefix_len) + branch_prefix_len = layout.prefix_len if replay_prefix else layout.prefix_len - fork_boundary + branches = pack_mamba_branches( + projected, + prefix_len=layout.prefix_len, + tail_len=branch_prefix_len, + completion_lens=tuple(physical_completion_lens), + ) + + prefix_y = projected.new_empty((0, 1, mixer.cp.d_inner_local_tpcp)) + if replay_prefix: + branch_y, _, _, _ = _scan_mamba_projected_segment(mixer, branches, include_gate=False) + else: + conv_context = None + final_state = None + if fork_boundary: + prefix_y, _, conv_context, final_state = _scan_mamba_projected_segment( + mixer, projected[:fork_boundary], capture_state=True, include_gate=False + ) + if conv_context is None or final_state is None: + raise RuntimeError( + "shared-prefix CP Mamba scan did not return a differentiable state" + ) + branch_y, _, _, _ = _scan_mamba_projected_segment( + mixer, + branches, + conv_context=conv_context, + ssm_initial_state=final_state, + include_gate=False, + ) + return merge_mamba_branches( + prefix_y, + branch_y, + tail_len=branch_prefix_len, + completion_lens=tuple(physical_completion_lens), + ) + + +def _forward_mamba_layer_shared_prefix_cp_impl( + layer: MambaLayer, + hidden_states: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, + *, + replay_prefix: bool, + packed_recurrence: bool = False, + ragged_state_fork: bool = False, +) -> Tensor: + """Run a TP/SP/CP Mamba star with state-fork or uninterrupted-replay prefix handling. + + The input projection first gathers sequence-parallel TP shards, then the CP adapter converts + local zigzag sequence shards into one canonical global sequence with TP/CP-local channels. + State-fork scans the chunk-aligned prefix head once, expands its differentiable state, and + replays only the unaligned prompt tail per branch. Replay is retained as a parity baseline and + repeats the full projected prefix per branch. The packed result returns through the inverse CP + transform and TP output-projection reduce-scatter. + """ + mixer = layer.mixer + _validate_mamba_fork(mixer) + cp = mixer.cp + + residual = hidden_states.float() if layer.config.fp32_residual_connection else hidden_states + normalized = apply_module(layer.norm)(hidden_states.to(dtype=layer.config.params_dtype)) + projected, _ = mixer.in_proj(normalized) + # Input SP gather: [physical/(TP*C), 1, hidden] -> [physical/C, 1, projection/TP]. + # Gated RMSNorm runs after the inverse CP transform, in this same local + # token/channel order. The scan does not use z, so keep it here instead of + # exchanging and packing it through every branch. Copy to avoid retaining + # the entire projection allocation solely for the gate. + local_z = projected[..., : cp.d_inner_local_tp].contiguous() + projected = cp.pre_conv_ssm(projected, include_gate=False) + physical_len = projected.shape[0] + if physical_len < layout.total_len: + raise ValueError( + f"parallel physical length {physical_len} is shorter than shared layout " + f"{layout.total_len}" + ) + + if ragged_state_fork: + if replay_prefix or packed_recurrence: + raise ValueError("Ragged state forking cannot use prefix replay") + from megatron.core.ssm.mamba_ragged import scan_mamba_ragged_forest + + packed_y = scan_mamba_ragged_forest(mixer, projected, layout) + elif packed_recurrence: + if replay_prefix: + raise ValueError("Packed recurrence and per-root replay are separate modes") + # Keep optional convolution/scan dependencies lazy for non-Mamba models. + from megatron.core.ssm.mamba_forest_replay import scan_mamba_forest_replay + + packed_y = scan_mamba_forest_replay(mixer, projected, layout) + else: + pieces = [] + for offset, root in layout.iter_roots(): + end = offset + root.total_len + if end == layout.total_len: + end = physical_len + pieces.append( + _scan_mamba_shared_prefix_root( + mixer, projected[offset:end], root, replay_prefix=replay_prefix + ) + ) + packed_y = pieces[0] if len(pieces) == 1 else torch.cat(pieces, dim=0) + # Canonical channel shards -> local zigzag sequence shards with full channels. + packed_y = cp.post_conv_ssm(packed_y) + packed_y = mixer.norm(packed_y, local_z) + packed_output, output_bias = mixer.out_proj(packed_y) + + with layer.bias_dropout_add_exec_handler(): + return layer.mamba_bda(training=layer.training, fused=layer.config.bias_dropout_fusion)( + (packed_output, output_bias), residual, layer.hidden_dropout + ) + + +def _forward_mamba_layer_shared_prefix_cp_state_fork( + layer: MambaLayer, hidden_states: Tensor, layout: SharedPrefixLayout | SharedPrefixForestLayout +) -> Tensor: + """Optimized CP Mamba: fork at a scan-chunk boundary and replay the prompt tail.""" + return _forward_mamba_layer_shared_prefix_cp_impl( + layer, hidden_states, layout, replay_prefix=False + ) + + +def _forward_mamba_layer_shared_prefix_cp_replay( + layer: MambaLayer, hidden_states: Tensor, layout: SharedPrefixLayout +) -> Tensor: + """Exact CP Mamba fallback: replay each prefix without an explicit state boundary.""" + return _forward_mamba_layer_shared_prefix_cp_impl( + layer, hidden_states, layout, replay_prefix=True + ) + + +def _forward_mamba_layer_shared_prefix_cp_packed_fused_oracle( + layer: MambaLayer, hidden_states: Tensor, layout: SharedPrefixLayout +) -> Tensor: + """Diagnostic dense replay through Mamba's unchanged packed fused path. + + This deliberately gives up Mamba prefix sharing. It reconstructs the ordinary + branch-major packed batch from the canonical star, invokes ``MambaLayer.forward`` + with the same ``PackedSeqParams`` contract as dense NeMo-RL sequence packing, and + then folds the result back to one shared-prefix star. Consequently the oracle + exercises ``mamba_split_conv1d_scan_combined(seq_idx=...)`` rather than the + decomposed causal-convolution/chunk-scan implementation used by the optimized + state-fork path. + + The helper is a correctness fallback and an isolation oracle. It is selected + explicitly with ``NRL_SP_MAMBA_IMPL=packed_fused``; the optimized state-fork + implementation remains the default until end-to-end GPU parity validates a + safer default. + """ + mixer = layer.mixer + if not isinstance(mixer, MambaMixer): + raise TypeError("shared-prefix packed-fused Mamba oracle requires MambaMixer") + tp_group = mixer.pg_collection.tp + cp_group = mixer.pg_collection.cp + tp_size = tp_group.size() + cp_size = cp_group.size() + physical_len = hidden_states.shape[0] * tp_size * cp_size + _validate_shared_prefix_physical_length( + layout, + physical_len, + tp_size=tp_size, + cp_size=cp_size, + sequence_parallel=bool(mixer.config.sequence_parallel), + ) + + # [star/(TP*CP),1,H] -> one canonical global star. The oracle runs under + # no_grad, but use the same collectives/order as the differentiable MTP + # reconstruction so the forward topology is representative. + cp_local_star = hidden_states + if tp_size > 1: + cp_local_star = tensor_parallel.gather_from_sequence_parallel_region( + cp_local_star, tensor_parallel_output_grad=False, group=tp_group + ) + if cp_size > 1: + rank_order_star = tensor_parallel.gather_from_sequence_parallel_region( + cp_local_star, tensor_parallel_output_grad=True, group=cp_group + ) + rank_order_indices = torch.cat( + [ + layout.cp_local_indices(physical_len, cp_size, rank, hidden_states.device) + for rank in range(cp_size) + ] + ) + inverse_order = torch.empty_like(rank_order_indices) + inverse_order[rank_order_indices] = torch.arange( + physical_len, device=hidden_states.device, dtype=torch.long + ) + global_star = rank_order_star.index_select(0, inverse_order) + else: + global_star = cp_local_star + if global_star.shape != (physical_len, 1, hidden_states.shape[-1]): + raise RuntimeError( + "shared-prefix packed-fused Mamba oracle reconstructed an invalid star shape" + ) + + # Drop topology-only padding after the final completion. Dense NeMo-RL + # padding is per branch, and those physical branch tails are already part of + # layout.completion_lens. + branch_indices = layout.dense_branch_indices(hidden_states.device) + branch_lengths = [int(indices.numel()) for indices in branch_indices] + dense_global = torch.cat( + [global_star.index_select(0, indices) for indices in branch_indices], dim=0 + ) + dense_total = sum(branch_lengths) + if dense_global.shape[0] != dense_total: + raise RuntimeError("shared-prefix packed-fused Mamba oracle built an invalid dense batch") + if cp_size > 1 and any(length % (2 * cp_size) for length in branch_lengths): + raise ValueError( + "shared-prefix packed-fused Mamba oracle requires each dense branch " + "to divide the CP zigzag quantum" + ) + + # Match NeMo-RL's per-sequence CP zigzag packing, then its TP sequence + # parallel shard. This ordering is distinct from the whole-star CP layout. + if cp_size > 1: + dense_cp_indices = [] + branch_offset = 0 + for branch_length in branch_lengths: + dense_cp_indices.append( + branch_offset + + layout.cp_local_indices( + branch_length, cp_size, cp_group.rank(), hidden_states.device + ) + ) + branch_offset += branch_length + dense_cp_indices = torch.cat(dense_cp_indices) + dense_cp_local = dense_global.index_select(0, dense_cp_indices) + else: + dense_cp_local = dense_global + dense_local = dense_cp_local + if tp_size > 1: + dense_local = tensor_parallel.scatter_to_sequence_parallel_region( + dense_local, group=tp_group + ) + + cumulative_lengths = [0] + for branch_length in branch_lengths: + cumulative_lengths.append(cumulative_lengths[-1] + branch_length) + cu_seqlens = torch.tensor(cumulative_lengths, device=hidden_states.device, dtype=torch.int32) + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens, + cu_seqlens_kv_padded=cu_seqlens, + max_seqlen_q=max(branch_lengths), + max_seqlen_kv=max(branch_lengths), + total_tokens=dense_total, + ) + dense_output_local = layer( + hidden_states=dense_local, attention_mask=None, packed_seq_params=packed_seq_params + ) + if isinstance(dense_output_local, tuple): + dense_output_local = dense_output_local[0] + + dense_output_cp_local = dense_output_local + if tp_size > 1: + dense_output_cp_local = tensor_parallel.gather_from_sequence_parallel_region( + dense_output_cp_local, tensor_parallel_output_grad=False, group=tp_group + ) + if cp_size > 1: + dense_rank_order = tensor_parallel.gather_from_sequence_parallel_region( + dense_output_cp_local, tensor_parallel_output_grad=True, group=cp_group + ) + rank_order_dense_indices = [] + for rank in range(cp_size): + branch_offset = 0 + for branch_length in branch_lengths: + rank_order_dense_indices.append( + branch_offset + + layout.cp_local_indices(branch_length, cp_size, rank, hidden_states.device) + ) + branch_offset += branch_length + rank_order_dense_indices = torch.cat(rank_order_dense_indices) + inverse_dense_order = torch.empty_like(rank_order_dense_indices) + inverse_dense_order[rank_order_dense_indices] = torch.arange( + dense_total, device=hidden_states.device, dtype=torch.long + ) + dense_output_global = dense_rank_order.index_select(0, inverse_dense_order) + else: + dense_output_global = dense_output_cp_local + + # One prompt output is sufficient: every dense branch has the same causal + # prompt and the oracle is forward-only. Completion outputs remain tied to + # their corresponding branch. + prefix_output = dense_output_global[: layout.prefix_len] + completion_outputs = [] + branch_offset = 0 + for completion_length, branch_length in zip( + layout.completion_lens, branch_lengths, strict=True + ): + completion_outputs.append( + dense_output_global[ + branch_offset + + layout.prefix_len : branch_offset + + layout.prefix_len + + completion_length + ] + ) + branch_offset += branch_length + star_output_global = torch.cat([prefix_output, *completion_outputs], dim=0) + trailing_padding = physical_len - layout.total_len + if trailing_padding: + star_output_global = torch.cat( + [ + star_output_global, + torch.zeros( + trailing_padding, + 1, + star_output_global.shape[-1], + dtype=star_output_global.dtype, + device=star_output_global.device, + ), + ], + dim=0, + ) + if cp_size > 1: + star_output_cp_local = star_output_global.index_select( + 0, layout.cp_local_indices(physical_len, cp_size, cp_group.rank(), hidden_states.device) + ) + else: + star_output_cp_local = star_output_global + if tp_size > 1: + return tensor_parallel.scatter_to_sequence_parallel_region( + star_output_cp_local, group=tp_group + ) + return star_output_cp_local + + +def _forward_mamba_layer_shared_prefix_cp( + layer: MambaLayer, hidden_states: Tensor, layout: SharedPrefixLayout | SharedPrefixForestLayout +) -> Tensor: + """Select the optimized Mamba path or correctness-first packed fallback.""" + implementation = os.environ.get("NRL_SP_MAMBA_IMPL", "state_fork") + if implementation == "state_fork": + return _forward_mamba_layer_shared_prefix_cp_state_fork(layer, hidden_states, layout) + if implementation in ("replay_prefix_training", "replay_prefix"): + # A training experiment: one uninterrupted scan per root can reduce + # launch/backward overhead, at the cost of repeating the recurrent + # prefix. Projection, attention and MoE still execute the shared layout. + # Training mode retains state forking during evaluation. The explicit + # replay_prefix mode also runs replay in evaluation, allowing numerical + # logprob checks to exercise the changed training recurrence itself. + return _forward_mamba_layer_shared_prefix_cp_impl( + layer, + hidden_states, + layout, + replay_prefix=implementation == "replay_prefix" or layer.training, + ) + if implementation in ("packed_recurrence_training", "packed_recurrence"): + return _forward_mamba_layer_shared_prefix_cp_impl( + layer, + hidden_states, + layout, + replay_prefix=False, + packed_recurrence=implementation == "packed_recurrence" or layer.training, + ) + if implementation == "packed_fused": + if isinstance(layout, SharedPrefixForestLayout): + raise NotImplementedError( + "packed_fused diagnostic supports one root; forests require state_fork" + ) + return _forward_mamba_layer_shared_prefix_cp_packed_fused_oracle( + layer, hidden_states, layout + ) + if implementation in ("ragged_state_fork_training", "ragged_state_fork"): + return _forward_mamba_layer_shared_prefix_cp_impl( + layer, + hidden_states, + layout, + replay_prefix=False, + ragged_state_fork=implementation == "ragged_state_fork" or layer.training, + ) + raise ValueError( + "NRL_SP_MAMBA_IMPL must be 'state_fork', 'replay_prefix_training', " + "'replay_prefix', 'packed_recurrence_training', 'packed_recurrence', 'packed_fused', " + "'ragged_state_fork_training' or 'ragged_state_fork', " + f"got {implementation!r}" + ) + + +def _has_nonzero_config_value(value) -> bool: + """Return whether a scalar or per-layer configuration contains a nonzero value.""" + if value is None: + return False + if isinstance(value, (list, tuple)): + return any(float(item) != 0.0 for item in value) + return float(value) != 0.0 + + +def _validate_shared_prefix_physical_length( + layout: SharedPrefixLayout | SharedPrefixForestLayout, + physical_len: int, + *, + tp_size: int, + cp_size: int, + sequence_parallel: bool, +) -> None: + """Validate the global star length against its negotiated topology/padding contract.""" + physical_len = int(physical_len) + if tp_size > 1: + topology_multiple = 2 * tp_size * cp_size + elif cp_size > 1: + topology_multiple = 2 * cp_size + else: + topology_multiple = 1 + if isinstance(layout, SharedPrefixForestLayout): + multiple = layout.padding_multiple or topology_multiple + if physical_len % multiple or not 0 <= physical_len - layout.total_len < multiple: + raise ValueError("shared-prefix forest must use minimal global topology padding") + for root in layout.roots: + _validate_shared_prefix_physical_length( + root, + ((root.total_len + multiple - 1) // multiple) * multiple, + tp_size=tp_size, + cp_size=cp_size, + sequence_parallel=sequence_parallel, + ) + return + padding = physical_len - layout.total_len + + if layout.padding_multiple is not None: + padding_multiple = layout.padding_multiple + if padding_multiple % topology_multiple: + raise ValueError( + "shared-prefix padding_multiple must be divisible by the topology quantum: " + f"M={padding_multiple}, Q={topology_multiple}, TP={tp_size}, CP={cp_size}" + ) + for branch, (physical_completion, logical_completion) in enumerate( + zip(layout.completion_lens, layout.logical_completion_lens, strict=True) + ): + if (layout.prefix_len + physical_completion) % padding_multiple or not ( + 0 <= physical_completion - logical_completion < padding_multiple + ): + raise ValueError( + "shared-prefix physical completion span must use the minimal per-branch " + "padding to padding_multiple: " + f"branch={branch}, prefix={layout.prefix_len}, " + f"logical={logical_completion}, physical={physical_completion}, " + f"M={padding_multiple}" + ) + if physical_len % padding_multiple: + raise ValueError( + "shared-prefix physical length must be divisible by padding_multiple: " + f"physical={physical_len}, M={padding_multiple}" + ) + if not 0 <= padding < padding_multiple: + raise ValueError( + "shared-prefix input must use the minimal trailing pad to padding_multiple: " + f"physical={physical_len}, layout={layout.total_len}, M={padding_multiple}" + ) + return + + if tp_size > 1: + if physical_len % topology_multiple: + raise ValueError( + "shared-prefix TP/SP physical length must be divisible by 2 * tensor parallel " + "size * context parallel size" + ) + if not 0 <= padding < topology_multiple: + raise ValueError( + "shared-prefix TP/SP input must use the minimal trailing pad to a 2*TP*CP " + f"multiple: physical={physical_len}, layout={layout.total_len}, " + f"TP={tp_size}, CP={cp_size}" + ) + elif cp_size == 1: + if physical_len != layout.total_len: + raise ValueError( + f"packed sequence length {physical_len} does not match layout {layout.total_len}" + ) + else: + if physical_len % topology_multiple: + raise ValueError( + "shared-prefix CP physical length must be divisible by 2 * context parallel size" + ) + if not 0 <= padding < topology_multiple: + raise ValueError( + "shared-prefix CP input must use the minimal trailing pad to a 2*CP multiple: " + f"physical={physical_len}, layout={layout.total_len}, CP={cp_size}" + ) + + +def _validate_hybrid_stack( + stack, hidden_states: Tensor, layout: SharedPrefixLayout | SharedPrefixForestLayout +) -> None: + # New upstream features must not silently bypass their ordinary execution. + # These combinations need dedicated shared-layout integrations and GPU tests. + if getattr(stack.config, "moe_num_hash_layers", 0): + raise NotImplementedError("shared-prefix Hybrid does not yet support hash MoE routing") + if getattr(stack.config, "quant_recipe", None) is not None: + raise NotImplementedError("shared-prefix Hybrid does not yet support quantization recipes") + if getattr(stack.config, "wide_residual", None) is not None: + raise NotImplementedError("shared-prefix Hybrid does not yet support wide residual streams") + if getattr(stack.config, "enable_mhc_connections", False): + raise NotImplementedError("shared-prefix Hybrid does not yet support mHC connections") + if getattr(stack.config, "attn_logit_softcapping", None) is not None: + raise NotImplementedError("shared-prefix attention does not yet support logit softcapping") + tp_size = stack.tp_group.size() + sequence_parallel = bool(stack.config.sequence_parallel) + if tp_size > 1 and not sequence_parallel: + raise NotImplementedError("shared-prefix Hybrid TP>1 requires sequence parallelism") + if tp_size == 1 and sequence_parallel: + raise NotImplementedError("shared-prefix Hybrid sequence parallelism requires TP>1") + if stack.config.tensor_model_parallel_size != tp_size: + raise RuntimeError( + "shared-prefix Hybrid tensor-parallel config does not match its process group" + ) + if stack.pg_collection.tp.size() != tp_size: + raise RuntimeError("shared-prefix Hybrid tensor-parallel groups disagree on their size") + if stack.pp_group.size() != 1: + raise NotImplementedError("shared-prefix Hybrid adapter currently supports PP1 only") + cp_group = stack.pg_collection.cp + cp_size = cp_group.size() + if stack.config.context_parallel_size != cp_size: + raise RuntimeError( + "shared-prefix Hybrid context-parallel config does not match its process group" + ) + if stack.config.recompute_granularity == "full": + if stack.config.recompute_method != "uniform": + raise NotImplementedError( + 'shared-prefix Hybrid full recomputation currently supports only the uniform ' + 'method' + ) + if not isinstance(stack.config.recompute_num_layers, int) or ( + stack.config.recompute_num_layers < 1 + ): + raise ValueError( + "shared-prefix Hybrid uniform recomputation requires recompute_num_layers >= 1" + ) + if stack.config.fine_grained_activation_offloading: + raise NotImplementedError( + "shared-prefix Hybrid adapter does not support fine-grained activation offloading" + ) + if stack.config.cuda_graph_impl != "none": + raise NotImplementedError("shared-prefix Hybrid adapter does not support CUDA graphs") + if stack.config.fp8 or stack.config.fp4: + raise NotImplementedError("shared-prefix Hybrid adapter currently supports fp16/bf16 only") + if hidden_states.dtype not in (torch.float16, torch.bfloat16): + raise TypeError("fused shared-prefix attention requires fp16 or bf16 hidden states") + if hidden_states.ndim != 3 or hidden_states.shape[1] != 1: + raise ValueError( + "shared-prefix Hybrid input must have shape [sequence/(TP*CP), 1, hidden] " + "when sequence parallelism is enabled" + ) + sequence_shards = tp_size if sequence_parallel else 1 + physical_len = hidden_states.shape[0] * cp_size * sequence_shards + _validate_shared_prefix_physical_length( + layout, physical_len, tp_size=tp_size, cp_size=cp_size, sequence_parallel=sequence_parallel + ) + if stack.config.attention_dropout != 0.0 or stack.config.hidden_dropout != 0.0: + raise NotImplementedError("shared-prefix Hybrid adapter currently requires zero dropout") + if stack.config.window_size not in (None, (-1, -1)): + raise NotImplementedError( + "shared-prefix Hybrid adapter does not support sliding-window attention" + ) + if stack.config.softmax_type != "vanilla": + raise NotImplementedError("shared-prefix fused attention supports only vanilla softmax") + + num_moe_experts = stack.config.num_moe_experts + expert_bias_enabled = bool(getattr(stack.config, "moe_router_enable_expert_bias", False)) + if num_moe_experts is not None and num_moe_experts > 0: + if stack.config.moe_router_force_load_balancing: + raise NotImplementedError( + "shared-prefix Hybrid adapter does not support randomized forced MoE routing" + ) + if stack.config.moe_router_force_biased is not None: + raise NotImplementedError( + "shared-prefix Hybrid adapter does not support randomized forced MoE router bias" + ) + load_balancing = stack.config.moe_router_load_balancing_type + load_balancing_types = ( + [load_balancing] if isinstance(load_balancing, str) else load_balancing + ) + if any(item != "none" for item in load_balancing_types): + raise NotImplementedError( + "shared-prefix Hybrid adapter requires MoE router load balancing type 'none'" + ) + if _has_nonzero_config_value(stack.config.moe_aux_loss_coeff): + raise NotImplementedError( + "shared-prefix Hybrid adapter does not support MoE auxiliary router loss" + ) + if _has_nonzero_config_value(stack.config.moe_z_loss_coeff): + raise NotImplementedError( + "shared-prefix Hybrid adapter does not support MoE router z-loss" + ) + if _has_nonzero_config_value(stack.config.moe_input_jitter_eps): + raise NotImplementedError( + "shared-prefix Hybrid adapter does not support MoE input jitter" + ) + if expert_bias_enabled and any( + root.logical_completion_lens is None for _, root in layout.iter_roots() + ): + raise NotImplementedError( + "shared-prefix MoE expert-bias accounting requires explicit physical " + "branch padding and logical completion lengths" + ) + if stack.config.moe_expert_capacity_factor is not None: + raise NotImplementedError( + 'shared-prefix Hybrid adapter does not support MoE expert capacity or token ' + 'dropping' + ) + if getattr(stack.config, "mlp_chunks_for_training", 1) != 1: + raise NotImplementedError( + "shared-prefix Hybrid MoE adapter does not support training MLP chunking" + ) + + for layer in stack.layers: + if isinstance(layer, MambaLayer): + if not isinstance(layer.mixer, MambaMixer): + raise NotImplementedError( + "shared-prefix state forking only supports MambaMixer-backed Mamba layers" + ) + _validate_mamba_fork(layer.mixer) + if layer.mixer.pg_collection.tp.size() != tp_size: + raise RuntimeError( + "shared-prefix Mamba TP helper does not match the Hybrid stack TP group" + ) + if layer.mixer.cp.cp_size != cp_size: + raise RuntimeError( + "shared-prefix Mamba CP helper does not match the Hybrid stack CP group" + ) + elif isinstance(layer, TransformerLayer): + if not isinstance(layer.self_attention, (IdentityOp, SelfAttention)): + raise NotImplementedError( + 'shared-prefix Hybrid adapter supports standard self-attention and MLP/MoE ' + 'layers' + ) + if ( + isinstance(layer.self_attention, SelfAttention) + and layer.self_attention.checkpoint_core_attention + ): + raise NotImplementedError( + 'shared-prefix fused attention does not yet support selective ' + 'core-attention recomputation' + ) + if isinstance(layer.self_attention, SelfAttention) and ( + stack.config.qk_clip or stack.config.log_max_attention_logit + ): + raise NotImplementedError( + 'shared-prefix fused attention does not yet produce QK-clipping/max-logit ' + 'statistics' + ) + if isinstance(layer.self_attention, SelfAttention): + if layer.self_attention.pg_collection.tp.size() != tp_size: + raise RuntimeError( + 'shared-prefix attention TP helper does not match the Hybrid stack TP ' + 'group' + ) + if layer.self_attention.pg_collection.cp.size() != cp_size: + raise RuntimeError( + 'shared-prefix attention CP helper does not match the Hybrid stack CP ' + 'group' + ) + if cp_size > 1 and isinstance(layer.self_attention, SelfAttention): + from megatron.core.models.hybrid.shared_prefix_fused import ( + _cp_kv_head_slices_for_destinations, + ) + + if stack.config.num_attention_heads % tp_size: + raise NotImplementedError( + "shared-prefix attention requires query heads divisible by TP size" + ) + query_heads = stack.config.num_attention_heads // tp_size + # This is the actual local K/V tensor width. When global KV heads are fewer than + # TP ranks, SelfAttention replicates one KV head across the relevant TP ranks. + kv_heads = layer.self_attention.num_query_groups_per_partition + _cp_kv_head_slices_for_destinations(query_heads, kv_heads, cp_size) + if expert_bias_enabled and getattr(layer, "is_moe_layer", False): + from megatron.core.transformer.moe.moe_layer import MoELayer + from megatron.core.transformer.moe.router import TopKRouter + + if not isinstance(layer.mlp, MoELayer) or not isinstance( + layer.mlp.router, TopKRouter + ): + raise NotImplementedError( + "shared-prefix expert-bias accounting requires an MCore " + "MoELayer with TopKRouter" + ) + else: + raise NotImplementedError( + f"shared-prefix state forking is not implemented for {type(layer).__name__}" + ) + + +def forward_hybrid_stack_shared_prefix( + stack, + hidden_states: Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, + *, + rotary_pos_emb: Tensor | tuple[Tensor, Tensor] | None = None, + position_embedding_type: str = "rope", +) -> Tensor: + """Explicit exact-prompt star forward for a supported ``HybridStack`` topology. + + Normal ``HybridStack.forward`` and ``Attention.forward`` behavior is unchanged unless this + function installs its scoped private forest descriptor. The descriptor is always removed in a + ``finally`` block, including when a layer raises. + """ + _validate_hybrid_stack(stack, hidden_states, layout) + has_attention = any( + isinstance(layer, TransformerLayer) and isinstance(layer.self_attention, SelfAttention) + for layer in stack.layers + ) + if position_embedding_type not in ("rope", "none"): + raise NotImplementedError( + "shared-prefix attention supports only RoPE or positionless Hybrid models" + ) + if has_attention and position_embedding_type == "rope" and rotary_pos_emb is None: + raise ValueError("position-aware rotary_pos_emb is required for shared-prefix attention") + if position_embedding_type == "none" and rotary_pos_emb is not None: + raise ValueError("positionless shared-prefix attention must not receive rotary_pos_emb") + + cp_group = stack.pg_collection.cp + tp_group = stack.pg_collection.tp + tp_size = tp_group.size() + sequence_shards = tp_size if stack.config.sequence_parallel else 1 + physical_len = hidden_states.shape[0] * cp_group.size() * sequence_shards + token_multiplicities = None + expert_bias_enabled = bool(getattr(stack.config, "moe_router_enable_expert_bias", False)) + if expert_bias_enabled: + token_multiplicities = layout.padded_token_multiplicities( + physical_len, + hidden_states.device, + # Match NeMo's dense get_packed_seq_padding_mask convention. + exclude_sequence_padding=( + getattr(stack.config, "moe_token_dispatcher_type", None) == "flex" + and getattr(stack.config, "moe_flex_dispatcher_backend", None) == "hybridep" + ), + ) + if cp_group.size() > 1: + token_multiplicities = token_multiplicities.index_select( + 0, + layout.cp_local_indices( + physical_len, cp_group.size(), cp_group.rank(), hidden_states.device + ), + ) + if tp_size > 1: + if token_multiplicities.numel() % tp_size: + raise RuntimeError( + "shared-prefix CP-local token multiplicities must divide evenly over TP" + ) + token_multiplicities = torch.chunk(token_multiplicities, tp_size, dim=0)[ + tp_group.rank() + ].contiguous() + if token_multiplicities.numel() != hidden_states.shape[0]: + raise RuntimeError( + "shared-prefix token multiplicity ownership does not match the local hidden rows" + ) + + @router_gating_token_blocks() + def forward_layer_range(hidden_states: Tensor, start: int, end: int) -> Tensor: + for layer in stack.layers[start:end]: + moe_layer = ( + layer.mlp if expert_bias_enabled and getattr(layer, "is_moe_layer", False) else None + ) + if moe_layer is not None: + if getattr(moe_layer, "_shared_prefix_token_multiplicities", None) is not None: + raise RuntimeError("nested shared-prefix MoE dispatch is not supported") + moe_layer._shared_prefix_token_multiplicities = token_multiplicities + try: + if isinstance(layer, MambaLayer): + if ( + cp_group.size() > 1 + or tp_size > 1 + or isinstance(layout, SharedPrefixForestLayout) + ): + hidden_states = _forward_mamba_layer_shared_prefix_cp( + layer, hidden_states, layout + ) + else: + hidden_states = _forward_mamba_layer_shared_prefix( + layer, hidden_states, layout + ) + elif isinstance(layer.self_attention, IdentityOp): + hidden_states = layer(hidden_states=hidden_states, attention_mask=None) + else: + attention = layer.self_attention + if getattr(attention, "_shared_prefix_forest", None) is not None: + raise RuntimeError( + "nested shared-prefix attention dispatch is not supported" + ) + attention._shared_prefix_forest = layout.forest + try: + hidden_states = layer( + hidden_states=hidden_states, + attention_mask=None, + rotary_pos_emb=rotary_pos_emb, + ) + finally: + del attention._shared_prefix_forest + finally: + if moe_layer is not None: + del moe_layer._shared_prefix_token_multiplicities + + if isinstance(hidden_states, tuple): + hidden_states = hidden_states[0] + return hidden_states + + if stack.config.recompute_granularity == "full" and stack.training: + chunk_size = stack.config.recompute_num_layers + for start in range(0, len(stack.layers), chunk_size): + end = min(start + chunk_size, len(stack.layers)) + + def custom_forward(hidden_states: Tensor, start=start, end=end) -> Tensor: + # Scopes live inside the callable so backward replay reconstructs them. + return forward_layer_range(hidden_states, start, end) + + hidden_states = tensor_parallel.checkpoint( + custom_forward, stack.config.distribute_saved_activations, hidden_states + ) + else: + hidden_states = forward_layer_range(hidden_states, 0, len(stack.layers)) + + if stack.post_process and stack.post_layer_norm: + hidden_states = stack.final_norm(hidden_states) + return make_viewless_tensor( + inp=hidden_states, requires_grad=hidden_states.requires_grad, keep_graph=True + ) + + +# The legacy scalar remains the validated CP1 negotiation surface. CP-aware integrations must +# negotiate against the collection so installing the CP track cannot regress existing CP1 runs. +# ``SHARED_PREFIX_CP_TRAINING_CAPABILITY`` is advertised independently so integrations can retain +# the CP1 fast path while requiring the validated CP>1 Hybrid forward/backward contract. +SHARED_PREFIX_TRAINING_CAPABILITY = "hybrid_star_cp1_tp1_v1" +SHARED_PREFIX_CP_TRAINING_CAPABILITY = "hybrid_star_cp_v1" +SHARED_PREFIX_EXPLICIT_PHYSICAL_PADDING_CAPABILITY = "hybrid_star_explicit_physical_padding_v1" +SHARED_PREFIX_MOE_EXPERT_BIAS_CAPABILITY = "hybrid_star_moe_expert_bias_v1" +SHARED_PREFIX_FULL_RECOMPUTE_CAPABILITY = "hybrid_star_full_uniform_recompute_v1" +SHARED_PREFIX_TP_SP_TRAINING_CAPABILITY = "hybrid_star_cp1_tp_sp_v1" +SHARED_PREFIX_CP_TP_SP_TRAINING_CAPABILITY = "hybrid_star_cp_tp_sp_v1" +# Validated target-model feature: Nemotron-H attention is positionless while +# Mamba remains state-positioned by sequence order. +SHARED_PREFIX_POSITIONLESS_ATTENTION_CAPABILITY = "hybrid_star_positionless_attention_v1" +# Validated MTP predictor feature: reconstruct dense attention/MLP or attention/MoE +# heads from the shared-prefix physical layout on the supported distributed TP/CP +# topologies. +SHARED_PREFIX_MTP_DENSE_HEADS_CAPABILITY = "hybrid_star_mtp_dense_heads_v1" +SHARED_PREFIX_GROUP_EXECUTION_CAPABILITY = "hybrid_forest_stable_router_v1" +# Topology and feature capabilities are independent so integrations can negotiate their exact +# validated conjunction without inferring support from a broader aggregate token. +SHARED_PREFIX_TRAINING_CAPABILITIES = frozenset( + { + SHARED_PREFIX_TRAINING_CAPABILITY, + SHARED_PREFIX_CP_TRAINING_CAPABILITY, + SHARED_PREFIX_EXPLICIT_PHYSICAL_PADDING_CAPABILITY, + SHARED_PREFIX_MOE_EXPERT_BIAS_CAPABILITY, + SHARED_PREFIX_FULL_RECOMPUTE_CAPABILITY, + SHARED_PREFIX_TP_SP_TRAINING_CAPABILITY, + SHARED_PREFIX_CP_TP_SP_TRAINING_CAPABILITY, + SHARED_PREFIX_POSITIONLESS_ATTENTION_CAPABILITY, + SHARED_PREFIX_MTP_DENSE_HEADS_CAPABILITY, + SHARED_PREFIX_GROUP_EXECUTION_CAPABILITY, + } +) diff --git a/megatron/core/models/hybrid/shared_prefix_fused.py b/megatron/core/models/hybrid/shared_prefix_fused.py new file mode 100644 index 00000000000..5756377f4b5 --- /dev/null +++ b/megatron/core/models/hybrid/shared_prefix_fused.py @@ -0,0 +1,1572 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Fused shared-prefix (tree/forest) attention — flash-composed passes with exact backward. + +Self-contained port of the optimized kernel suite from the aresk_shared_prefix sandbox (opt1–opt11 +on branch opt1-plancache-tritonmerge @ 3aa3faf03; bench/parity harness lives there under +examples/shared_prefix_attention). Public entry points: + +- ``flash_composed_forest_attention_fused(q, k, v, node_start, node_len, node_parent)`` — general + DFS-preorder forest attention as few flash varlen passes merged by online-softmax LSE, with the + EXACT backward (merged-output substitution: plain autograd drops the inter-pass normalizer + coupling because flash's LSE output carries no gradient). +- ``flash_composed_forest_attention(q, k, v, forest)`` — multi-group star forests + (``[(offset, prefix_len, completion_lens), ...]``), fused into one plan. + +Env knobs (defaults tuned on GB200): NRL_SP_CHAINFIRST (hybrid chain plan), NRL_SP_QSLICE +(zero-copy q views), NRL_SP_STREAMS (stream overlap), and NRL_SP_COMBINE (cross-pass +consolidation, default off). Experimental Triton KV gather, backward glue, and dQ assembly have +separate, default-off opt-ins so each can be parity-qualified independently. +NRL_SP_DETERMINISTIC_BACKWARD is a default-off diagnostic that selects FlashAttention's +deterministic backward. The retained Triton LSE merge is production-disabled pending a full-model +corruption fix. +GB200 net vs block-diagonal at equal work: star-like/balanced trees 1.59x training / 1.56x +logprob; deep branched trees 1.00x / 1.01x at a 1.10x FLOP ceiling (91-92% kernel efficiency). +""" + +import os +from contextlib import nullcontext +from itertools import pairwise +from typing import List + +import torch + +from megatron.core.tensor_parallel.mappings import all_to_all_hp2sp, all_to_all_sp2hp + +# --- Optimization: plan caching + fused (Triton) LSE merge ------------------------------------ +# The fused kernel's overhead vs raw flash is NOT attention math; it is (a) rebuilding the pass +# plan (pure-Python node loops + Python-int index lists -> H2D copies) on EVERY call -- the same +# bin layout recurs across all ~50 layers of a step -- and (b) the eager online-softmax merge, +# which upcasts every pass output to fp32 and round-trips full [total, np, hn] tensors through +# memory several times per forward. (a) is fixed by an LRU plan cache keyed on the node arrays; +# (b) by a single Triton kernel that reads each pass's output/LSE once and writes the merged +# output (+ final LSE) once, fp32 math in-register, bf16 out. The LSE merge is fail-closed below; +# other Triton optimizations have independent default-off gates and automatically fall back when +# Triton is unavailable. +try: + import triton + import triton.language as tl + + HAVE_TRITON = True +except ImportError: + HAVE_TRITON = False + + +# Default OFF in this production copy: the Triton merge kernel is bit-correct in isolation +# (kernel-vs-eager 3e-4 on captured in-model pass tensors) but nondeterministically corrupts +# the cross-pass rows when run inside the full HybridModel process (0.2-0.3 rel logits error, +# Heisenbug: any sync/instrumentation in the pass region masks it; streams/tile-config/ +# sync-before-merge all ruled out — suspected Triton runtime interaction with TE/mamba +# kernels in-process). The eager merge is stably exact in-model and the end-to-end cost is +# small (fwd 3.54x vs flex 3.11x with eager merge). Fail closed if a sandbox happens to export +# the old opt-in: known-corrupt execution must not silently enter a production training run. +def _resolve_fused_merge_setting(): + requested = os.environ.get("NRL_SP_FUSED_MERGE", "0") not in ("0", "", "false", "False") + if requested: + raise RuntimeError( + "NRL_SP_FUSED_MERGE is disabled in the production shared-prefix port: " + "the Triton merge has known nondeterministic full-HybridModel corruption" + ) + return False + + +def _resolve_experimental_triton_setting(env_name): + """Resolve a default-off opt-in for one independently parity-qualified Triton path. + + These kernels are intentionally experimental. Reject ambiguous values so a launcher typo + cannot silently turn one on before its target topology has passed numerical parity. + """ + value = os.environ.get(env_name, "0") + normalized = value.lower() + if normalized in ("", "0", "false"): + return False + if normalized in ("1", "true"): + return True + raise RuntimeError(f"{env_name} must be one of 0, 1, false, or true; got {value!r}") + + +def _resolve_deterministic_backward_setting(): + """Resolve the opt-in deterministic FlashAttention backward diagnostic. + + Reject unknown values instead of silently selecting the faster nondeterministic path when a + launcher misspells the diagnostic setting. + """ + value = os.environ.get("NRL_SP_DETERMINISTIC_BACKWARD", "0") + normalized = value.lower() + if normalized in ("0", "false"): + return False + if normalized in ("1", "true"): + return True + raise RuntimeError( + "NRL_SP_DETERMINISTIC_BACKWARD must be one of 0, 1, false, or true; " f"got {value!r}" + ) + + +_SP_FUSED_MERGE = _resolve_fused_merge_setting() +_SP_FUSED_KV_GATHER = _resolve_experimental_triton_setting("NRL_SP_FUSED_KV_GATHER") +_SP_FUSED_BACKWARD_GLUE = _resolve_experimental_triton_setting("NRL_SP_FUSED_BACKWARD_GLUE") +_SP_FUSED_DQ_ASSEMBLY = _resolve_experimental_triton_setting("NRL_SP_FUSED_DQ_ASSEMBLY") +_SP_DETERMINISTIC_BACKWARD = _resolve_deterministic_backward_setting() +# merge/dq-assembly kernel tile config (swept on GB200; override for other parts) +_SP_MERGE_BT = int(os.environ.get("NRL_SP_MERGE_BT", "16")) +_SP_MERGE_WARPS = int(os.environ.get("NRL_SP_MERGE_WARPS", "8")) + +_PLAN_CACHE: dict = {} +_PLAN_CACHE_MAX = 128 + + +if HAVE_TRITON: + + @triton.jit + def _sp_merge_fwd_kernel( + o0, + o1, + o2, + o3, + o4, + o5, + o6, # pass outputs, [rows_p, np, HN] (dtype of q) + l0, + l1, + l2, + l3, + l4, + l5, + l6, # pass LSEs, fp32 [np, rows_p] + i0, + i1, + i2, + i3, + i4, + i5, + i6, # int32 [total]: token -> row in pass (or -1) + r0, + r1, + r2, + r3, + r4, + r5, + r6, # rows_p per pass (for LSE stride) + out_ptr, # merged output [total, np, HN] (dtype of q) + lsef_ptr, # final LSE fp32 [np, total] + n_passes, + total, + np_: tl.constexpr, + HN: tl.constexpr, + BLOCK_T: tl.constexpr, + ): + # One program merges a BLOCK_T-token tile for one head: amortizes scheduling over + # 524k-programs-of-tiny-work (the v1 grid), keeps o loads coalesced along HN, and gives + # the compiler ILP across the tile. [BLOCK_T, HN] fp32 tile state. + tb = tl.program_id(0) + h = tl.program_id(1) + t = tb * BLOCK_T + tl.arange(0, BLOCK_T) + tmask = t < total + offs = tl.arange(0, HN) + m = tl.full([BLOCK_T], float("-inf"), dtype=tl.float32) + s = tl.zeros([BLOCK_T], dtype=tl.float32) + acc = tl.zeros([BLOCK_T, HN], dtype=tl.float32) + for p in tl.static_range(7): + if p < n_passes: + if p == 0: + idx_ptr, o_ptr, l_ptr, rows = i0, o0, l0, r0 + elif p == 1: + idx_ptr, o_ptr, l_ptr, rows = i1, o1, l1, r1 + elif p == 2: + idx_ptr, o_ptr, l_ptr, rows = i2, o2, l2, r2 + elif p == 3: + idx_ptr, o_ptr, l_ptr, rows = i3, o3, l3, r3 + elif p == 4: + idx_ptr, o_ptr, l_ptr, rows = i4, o4, l4, r4 + elif p == 5: + idx_ptr, o_ptr, l_ptr, rows = i5, o5, l5, r5 + else: + idx_ptr, o_ptr, l_ptr, rows = i6, o6, l6, r6 + r = tl.load(idx_ptr + t, mask=tmask, other=-1) + hit = (r >= 0) & tmask + lse = tl.load(l_ptr + h * rows + r, mask=hit, other=float("-inf")) + o = tl.load( + o_ptr + (r[:, None] * np_ + h) * HN + offs[None, :], + mask=hit[:, None], + other=0.0, + ).to(tl.float32) + m_new = tl.maximum(m, lse) + m_safe = tl.where(m_new == float("-inf"), 0.0, m_new) + scale_old = tl.where(m == float("-inf"), 0.0, tl.exp(m - m_safe)) + w = tl.where(hit, tl.exp(lse - m_safe), 0.0) + acc = acc * scale_old[:, None] + o * w[:, None] + s = s * scale_old + w + m = m_new + out = acc / tl.where(s == 0.0, 1.0, s)[:, None] + tl.store( + out_ptr + (t[:, None] * np_ + h) * HN + offs[None, :], + out.to(out_ptr.dtype.element_ty), + mask=tmask[:, None], + ) + tl.store(lsef_ptr + h * total + t, m + tl.log(s), mask=tmask) + + +if HAVE_TRITON: + + @triton.jit + def _sp_scale_gather_kernel( + do_ptr, # [total, np, HN] upstream grad (q dtype) + lse_ptr, # fp32 [np, rows] this pass's LSE + lsef_ptr, # fp32 [np, total] merged LSE + qidx_ptr, # int64 [rows] pass-row -> token (identity pass passes arange) + dox_ptr, # out [rows, np, HN] (q dtype): w * do[qidx] + rows, + total, + np_: tl.constexpr, + HN: tl.constexpr, + ): + rb = tl.program_id(0) + h = tl.program_id(1) + BLOCK_R: tl.constexpr = 16 + r = rb * BLOCK_R + tl.arange(0, BLOCK_R) + rmask = r < rows + offs = tl.arange(0, HN) + t = tl.load(qidx_ptr + r, mask=rmask, other=0) + lse_p = tl.load(lse_ptr + h * rows + r, mask=rmask, other=0.0) + w = tl.exp(lse_p - tl.load(lsef_ptr + h * total + t, mask=rmask, other=0.0)) + # zero-K padding rows (q-slice cross pass) report LSE=+inf: their weight must be 0, + # not inf, so their (zero) flash grads stay zero instead of turning NaN. + w = tl.where(lse_p > 1e30, 0.0, w) + do = tl.load( + do_ptr + (t[:, None] * np_ + h) * HN + offs[None, :], mask=rmask[:, None], other=0.0 + ).to(tl.float32) + tl.store( + dox_ptr + (r[:, None] * np_ + h) * HN + offs[None, :], + (do * w[:, None]).to(dox_ptr.dtype.element_ty), + mask=rmask[:, None], + ) + + @triton.jit + def _sp_scatter_accum_kernel( + dst_ptr, # fp32 [total, n, HN] accumulator + src_ptr, # [rows, n, HN] pass grad (q dtype) + idx_ptr, # int64 [rows] pass-row -> token + n_rows, + n_: tl.constexpr, + HN: tl.constexpr, + ): + rb = tl.program_id(0) + h = tl.program_id(1) + BLOCK_R: tl.constexpr = 16 + r = rb * BLOCK_R + tl.arange(0, BLOCK_R) + rmask = r < n_rows + offs = tl.arange(0, HN) + t = tl.load(idx_ptr + r, mask=rmask, other=0) + add = tl.load( + src_ptr + (r[:, None] * n_ + h) * HN + offs[None, :], mask=rmask[:, None], other=0.0 + ).to(tl.float32) + # atomic: with the consolidated cross pass a deep token owns one row PER ancestor level, + # so multiple programs may target the same destination row. + tl.atomic_add( + dst_ptr + (t[:, None] * n_ + h) * HN + offs[None, :], add, mask=rmask[:, None] + ) + + +if HAVE_TRITON: + + @triton.jit + def _sp_dq_merge_kernel( + d0, + d1, + d2, + d3, + d4, + d5, + d6, # per-pass dq contributions [rows_p, np, HN] (q dtype) + i0, + i1, + i2, + i3, + i4, + i5, + i6, # int32 [total]: token -> row in pass (or -1) + out_ptr, # final dq [total, np, HN] (q dtype) + n_passes, + total, + np_: tl.constexpr, + HN: tl.constexpr, + BLOCK_T: tl.constexpr, + ): + # opt10: dq final assembly as a gather-side sum over slots (mirror of the fwd merge, + # minus the LSE weighting — each pass's dqx is already its finished contribution). + # Replaces the fp32 [total, np, HN] accumulator + per-pass scatter/adds + final cast: + # every dqx is read once, dq written once, fp32 math in-register. + tb = tl.program_id(0) + h = tl.program_id(1) + t = tb * BLOCK_T + tl.arange(0, BLOCK_T) + tmask = t < total + offs = tl.arange(0, HN) + acc = tl.zeros([BLOCK_T, HN], dtype=tl.float32) + for p in tl.static_range(7): + if p < n_passes: + if p == 0: + idx_ptr, d_ptr = i0, d0 + elif p == 1: + idx_ptr, d_ptr = i1, d1 + elif p == 2: + idx_ptr, d_ptr = i2, d2 + elif p == 3: + idx_ptr, d_ptr = i3, d3 + elif p == 4: + idx_ptr, d_ptr = i4, d4 + elif p == 5: + idx_ptr, d_ptr = i5, d5 + else: + idx_ptr, d_ptr = i6, d6 + r = tl.load(idx_ptr + t, mask=tmask, other=-1) + hit = (r >= 0) & tmask + acc += tl.load( + d_ptr + (r[:, None] * np_ + h) * HN + offs[None, :], + mask=hit[:, None], + other=0.0, + ).to(tl.float32) + tl.store( + out_ptr + (t[:, None] * np_ + h) * HN + offs[None, :], + acc.to(out_ptr.dtype.element_ty), + mask=tmask[:, None], + ) + + @triton.jit + def _sp_gather_kv_kernel( + k_ptr, + v_ptr, # [total, ng, HN] sources (q dtype) + idx_ptr, # int64 [rows] pass-row -> token + kx_ptr, + vx_ptr, # [rows, ng, HN] destinations + k_stride_t, + k_stride_h, + k_stride_d, + v_stride_t, + v_stride_h, + v_stride_d, + ng: tl.constexpr, + HN: tl.constexpr, + ): + r = tl.program_id(0) + h = tl.program_id(1) + offs = tl.arange(0, HN) + t = tl.load(idx_ptr + r) + tl.store( + kx_ptr + (r * ng + h) * HN + offs, + tl.load(k_ptr + t * k_stride_t + h * k_stride_h + offs * k_stride_d), + ) + tl.store( + vx_ptr + (r * ng + h) * HN + offs, + tl.load(v_ptr + t * v_stride_t + h * v_stride_h + offs * v_stride_d), + ) + + +# Round-3: overlap the independent per-pass flash calls on side CUDA streams (they only join at +# the LSE merge / the gradient scatters), and gather K+V through one fused kernel into +# plan-cached workspace buffers (no per-call allocations, one index read for both tensors). +# NRL_SP_STREAMS=0 disables the stream overlap (kernels still fused). +_SP_STREAMS = os.environ.get("NRL_SP_STREAMS", "1") not in ("0", "", "false", "False") +# Consolidate all per-level cross passes into one flash call (any depth => 2 flash calls total). +# Requires the Triton merge path (the eager merge cannot handle a token owning multiple rows of +# one pass); plans are cached per effective mode so runtime flag flips stay correct. +_SP_COMBINE_CROSS = os.environ.get("NRL_SP_COMBINE", "1") not in ("0", "", "false", "False") +_SP_STREAM_POOL: List = [] +_SP_STREAM_POOL_N = 4 + + +def _sp_streams(): + if not _SP_STREAMS or not torch.cuda.is_available(): + return None + if not _SP_STREAM_POOL: + _SP_STREAM_POOL.extend(torch.cuda.Stream() for _ in range(_SP_STREAM_POOL_N)) + return _SP_STREAM_POOL + + +def _sp_fused_kv_gather_effective(): + return HAVE_TRITON and _SP_FUSED_KV_GATHER + + +def _sp_fused_backward_glue_effective(): + return HAVE_TRITON and _SP_FUSED_BACKWARD_GLUE + + +def _sp_fused_dq_assembly_effective(slot_pass): + # The kernel has seven statically-unrolled slot arguments. Deeper trees must retain the + # eager accumulator instead of tripping the kernel's assertion after an explicit opt-in. + return HAVE_TRITON and _SP_FUSED_DQ_ASSEMBLY and len(slot_pass) <= 7 + + +def _gather_kv(k, v, k_idx): + """Fused K+V gather (one index read, both tensors) via Triton. + + K/V commonly arrive as strided views of Megatron's interleaved mixed-QKV projection, so the + source strides must be explicit. Treating them as packed ``[total, ng, hn]`` tensors silently + reads neighboring Q/K fields on cross-attention passes. The outputs are newly allocated and + contiguous; the non-Triton path falls back to two stride-aware ``index_select`` calls. + """ + if _sp_fused_kv_gather_effective(): + rows = k_idx.numel() + ng, hn = k.shape[1], k.shape[2] + kx = torch.empty(rows, ng, hn, dtype=k.dtype, device=k.device) + vx = torch.empty(rows, ng, hn, dtype=v.dtype, device=v.device) + _sp_gather_kv_kernel[(rows, ng)]( + k, + v, + k_idx, + kx, + vx, + k.stride(0), + k.stride(1), + k.stride(2), + v.stride(0), + v.stride(1), + v.stride(2), + ng=ng, + HN=hn, + ) + return kx, vx + return k.index_select(0, k_idx), v.index_select(0, k_idx) + + +def _sp_combine_effective(): + return _SP_COMBINE_CROSS and HAVE_TRITON and _SP_FUSED_MERGE + + +def _plan_key(node_start, node_len, node_parent, device, *, full_context: bool = False): + return ( + tuple(int(x) for x in node_start), + tuple(int(x) for x in node_len), + tuple(int(x) for x in node_parent), + str(device), + _sp_combine_effective(), + _SP_CHAINFIRST, + _SP_QSLICE, + full_context, + ) + + +def _validate_forest_dfs_preorder(node_start, node_len, node_parent): + """Reject layouts whose node spans cannot represent a contiguous DFS-preorder forest.""" + ns = [int(x) for x in node_start] + nl = [int(x) for x in node_len] + par = [int(x) for x in node_parent] + if len(ns) != len(nl) or len(ns) != len(par): + raise ValueError("forest node_start, node_len, and node_parent must have equal lengths") + + cursor = 0 + for i, (start, length) in enumerate(zip(ns, nl)): + if start != cursor: + raise ValueError( + f"forest node {i} start={start} != expected {cursor} " + "(spans must be contiguous in array order)" + ) + if length <= 0: + raise ValueError(f"forest node {i} has non-positive length {length}") + cursor += length + + stack: List[int] = [] + for i, parent in enumerate(par): + if parent == -1: + stack = [i] + continue + while stack and stack[-1] != parent: + stack.pop() + if not stack or stack[-1] != parent: + raise ValueError( + f"forest/tree layout is not DFS-preorder at node {i} (parent {parent}): the fused " + "tree attention requires each node's subtree to be the contiguous run after it. " + "Emit nodes in DFS preorder (parent, then each child's full subtree)." + ) + stack.append(i) + + +def _forest_attention_plan_cached( + node_start, node_len, node_parent, device, *, full_context: bool = False +): + """Cached ``(total, passes, inv_maps)`` for a bin layout. The same layout is reused by every + attention layer of the step (and often across steps), so the Python plan construction and the + token->pass-row inverse maps (for the fused merge) are built once. ``inv_maps[p]`` is an int32 + ``[total]`` tensor mapping token -> its row in pass ``p`` (-1 if the token is not a query of + that pass; pass 0 -- the self pass -- is the identity).""" + _validate_forest_dfs_preorder(node_start, node_len, node_parent) + key = _plan_key(node_start, node_len, node_parent, device, full_context=full_context) + hit = _PLAN_CACHE.get(key) + if hit is not None: + return hit + plan = None + if full_context: + plan = _forest_attention_plan_full_context(node_start, node_len, node_parent, device) + elif _SP_CHAINFIRST: + plan = _forest_attention_plan_chainfirst(node_start, node_len, node_parent, device) + if plan is None: + plan = _forest_attention_plan(node_start, node_len, node_parent, device) + total, passes = plan + # never consolidate slice-form cross passes: combining would re-materialize the q gather + # (and its backward scatter) that the _QSlice views exist to avoid. + if ( + _sp_combine_effective() + and len(passes) > 2 + and not any(isinstance(p[0], _QSlice) for p in passes) + ): + # Consolidate every per-depth-level cross pass into ONE flash_varlen call: varlen just + # needs per-sequence contiguous q/k slices, and each (ancestor-span <- descendant-run) + # pair is one sequence regardless of which level it came from. Any tree then costs + # exactly 2 flash calls (self + cross) instead of depth+1 -- fewer launches, bigger + # kernels. The merge/backward kernels are unchanged: they operate per SLOT (a token's + # entry at one ancestor level), and cross slots simply share the combined pass's + # output/LSE tensors with different row maps. + self_pass = passes[0] + qpos_parts, kpos_parts, cuq, cuk = [], [], [0], [0] + level_row_start = [] # row offset of each level's block in the combined pass + for q_idx, k_idx, cu_q, cu_k, _mxq, _mxk, _c in passes[1:]: + level_row_start.append(cuq[-1]) + qpos_parts.append(q_idx) + kpos_parts.append(k_idx) + base_q, base_k = cuq[-1], cuk[-1] + cuq.extend((cu_q[1:].to(torch.long) + base_q).tolist()) + cuk.extend((cu_k[1:].to(torch.long) + base_k).tolist()) + qpos = torch.cat(qpos_parts) + kpos = torch.cat(kpos_parts) + mxq = max(cuq[i + 1] - cuq[i] for i in range(len(cuq) - 1)) + mxk = max(cuk[i + 1] - cuk[i] for i in range(len(cuk) - 1)) + combined = ( + qpos, + kpos, + torch.tensor(cuq, dtype=torch.int32, device=device), + torch.tensor(cuk, dtype=torch.int32, device=device), + mxq, + mxk, + False, + ) + # slots: slot 0 = self pass; slot j>=1 = the level-(j-1) rows INSIDE the combined pass. + slot_pass = [0] + [1] * (len(passes) - 1) + slot_inv = [torch.arange(total, dtype=torch.int32, device=device)] + for li, (q_idx, *_rest) in enumerate(passes[1:]): + inv = torch.full((total,), -1, dtype=torch.int32, device=device) + inv[q_idx] = ( + torch.arange(q_idx.numel(), dtype=torch.int32, device=device) + level_row_start[li] + ) + slot_inv.append(inv) + passes = [self_pass, combined] + identity = torch.arange(total, dtype=torch.long, device=device) + qidx64 = [identity, qpos] + entry = (total, passes, slot_inv, qidx64, slot_pass) + else: + inv_maps, qidx64 = [], [] + identity = torch.arange(total, dtype=torch.long, device=device) + for q_idx, *_rest in passes: + if q_idx is None: + inv = torch.arange(total, dtype=torch.int32, device=device) + qidx64.append(identity) + elif isinstance(q_idx, _QSlice): + # row r of the pass <-> token lo+r; gap tokens are padding rows, not queries. + inv = torch.full((total,), -1, dtype=torch.int32, device=device) + inv[q_idx.lo : q_idx.hi] = torch.arange( + q_idx.hi - q_idx.lo, dtype=torch.int32, device=device + ) + for g0, g1 in q_idx.gaps: + inv[g0:g1] = -1 + qidx64.append(torch.arange(q_idx.lo, q_idx.hi, dtype=torch.long, device=device)) + else: + inv = torch.full((total,), -1, dtype=torch.int32, device=device) + inv[q_idx] = torch.arange(q_idx.numel(), dtype=torch.int32, device=device) + qidx64.append(q_idx) + inv_maps.append(inv) + entry = (total, passes, inv_maps, qidx64, list(range(len(passes)))) + if len(_PLAN_CACHE) >= _PLAN_CACHE_MAX: + _PLAN_CACHE.pop(next(iter(_PLAN_CACHE))) + _PLAN_CACHE[key] = entry + return entry + + +def _merge_passes_triton(slot_pass, inv_maps, outs, lses, total, np_, hn, dtype, device): + """One-kernel online-softmax merge across SLOTS. A slot is one attended-set contribution for + a token (its own node, or one ancestor level); ``slot_pass[s]`` says which pass's output/LSE + tensors slot ``s`` reads (with the consolidated cross pass, every cross slot shares pass 1's + tensors under a different row map). Returns (o_merged [total, np, hn] in ``dtype``, lse_final + fp32 [np, total]).""" + MAXP = 7 + n = len(slot_pass) + assert n <= MAXP, f"fused merge supports <= {MAXP} slots (got {n}); deepen tl.static_range" + outs = [o.contiguous() for o in outs] + lses = [l.contiguous() for l in lses] + o_args, l_args, i_args, r_args = [], [], [], [] + for s in range(MAXP): + p = slot_pass[s] if s < n else slot_pass[0] + o_args.append(outs[p]) + l_args.append(lses[p]) + i_args.append(inv_maps[s] if s < n else inv_maps[0]) + r_args.append(lses[p].shape[1]) + out = torch.empty(total, np_, hn, dtype=dtype, device=device) + lse_final = torch.empty(np_, total, dtype=torch.float32, device=device) + bt, wp = _SP_MERGE_BT, _SP_MERGE_WARPS + _sp_merge_fwd_kernel[((total + bt - 1) // bt, np_)]( + *o_args, + *l_args, + *i_args, + *r_args, + out, + lse_final, + n, + total, + np_=np_, + HN=hn, + BLOCK_T=bt, + num_warps=wp, + ) + return out, lse_final + + +def _merge_dq_triton(slot_pass, inv_maps, dqxs, total, np_, hn, dtype, device): + """dq final assembly across slots: dq[t] = sum over slots s of dqxs[slot_pass[s]][inv_s[t]]. + One kernel, dqx tensors read once, dq written once in ``dtype`` (no fp32 accumulator).""" + MAXP = 7 + n = len(slot_pass) + assert n <= MAXP + d_args, i_args = [], [] + for s in range(MAXP): + p = slot_pass[s] if s < n else slot_pass[0] + d_args.append(dqxs[p]) + i_args.append(inv_maps[s] if s < n else inv_maps[0]) + dq = torch.empty(total, np_, hn, dtype=dtype, device=device) + bt, wp = _SP_MERGE_BT, _SP_MERGE_WARPS + _sp_dq_merge_kernel[((total + bt - 1) // bt, np_)]( + *d_args, *i_args, dq, n, total, np_=np_, HN=hn, BLOCK_T=bt, num_warps=wp + ) + return dq + + +# Chain-first plan (NRL_SP_CHAINFIRST=1, default on, auto-fallback): when the layout emits each +# node's continuation child immediately after it (chain-first DFS), maximal parent-adjacent runs +# ("chains") behave as plain CAUSAL sequences — a chain token's causal prefix within the run is +# exactly its in-chain ancestors. The self pass then uses per-CHAIN (not per-node) causal +# sequences, absorbing all within-chain cross attention (for spine-dominated trees that deletes +# most cross rows). Each chain with ancestors ABOVE its head attends one contiguous k-range +# [path_start, chain_start) — fat, flash-friendly K instead of per-level skinny spans — packed +# into a single non-causal cross pass. Falls back to the per-level plan when any cross-needing +# chain's ancestor range is non-contiguous (e.g. interior non-first children). +_SP_CHAINFIRST = os.environ.get("NRL_SP_CHAINFIRST", "1") not in ("0", "", "false", "False") + +# opt8: when the chain-first cross pass's query rows form one contiguous layout range (they do +# per tree: branches+siblings all sit after the spine), pass a zero-copy VIEW q[lo:hi] to flash +# instead of index_select (on branched_mc that gather+scatter round-trips ~176MB per fwd+bwd). +# Gaps between trees in multi-tree bins are covered by zero-length-K padding sequences: flash +# returns out=0 / dq=0 / LSE=+inf for those rows (probe-verified), the merge excludes them via +# inv=-1, and the backward scale kernel guards LSE=+inf -> weight 0. +_SP_QSLICE = os.environ.get("NRL_SP_QSLICE", "1") not in ("0", "", "false", "False") + + +class _QSlice: + """Marker for a cross pass whose q rows are the contiguous token range [lo, hi) (row r of + the pass <-> token lo+r), with ``gaps`` = token sub-ranges inside [lo, hi) that are only + zero-K padding sequences (not real queries of the pass).""" + + __slots__ = ("lo", "hi", "gaps") + + def __init__(self, lo, hi, gaps): + self.lo, self.hi, self.gaps = lo, hi, gaps + + def numel(self): + """Return the number of query rows in the contiguous slice.""" + return self.hi - self.lo + + +def _sel_rows(t, q_idx): + """Resolve a pass's q-row selector: None => identity, _QSlice => zero-copy view, + tensor => gather.""" + if q_idx is None: + return t + if isinstance(q_idx, _QSlice): + return t[q_idx.lo : q_idx.hi] + return t.index_select(0, q_idx) + + +def _forest_attention_plan_full_context(node_start, node_len, node_parent, device): + """Give each node its full ancestor KV path in one bottom-right causal pass. + + This trades repeated narrow GQA KV rows for removing partial-output rounding + and the forward LSE merge. It preserves the existing backward's KV scatter. + It does not emulate the original dense TE context-parallel arithmetic. + """ + starts = [int(x) for x in node_start] + lengths = [int(x) for x in node_len] + parents = [int(x) for x in node_parent] + kv_parts, cu_q, cu_k = [], [0], [0] + for node, length in enumerate(lengths): + if length == 0: + continue + path = [] + ancestor = node + while ancestor != -1: + path.append(ancestor) + ancestor = parents[ancestor] + path.reverse() + kv_parts.extend( + torch.arange(starts[i], starts[i] + lengths[i], device=device) for i in path + ) + cu_q.append(cu_q[-1] + length) + cu_k.append(cu_k[-1] + sum(lengths[i] for i in path)) + assert kv_parts, "empty forests must be handled before attention" + maximum_q = max(b - a for a, b in pairwise(cu_q)) + maximum_k = max(b - a for a, b in pairwise(cu_k)) + return cu_q[-1], [ + ( + None, + torch.cat(kv_parts), + torch.tensor(cu_q, dtype=torch.int32, device=device), + torch.tensor(cu_k, dtype=torch.int32, device=device), + maximum_q, + maximum_k, + True, + ) + ] + + +def _forest_attention_plan_chainfirst(node_start, node_len, node_parent, device): + """Hybrid chain decomposition (opt9; supersedes both the pure chain-first plan and the + per-level fallback). Always applicable: + + - SELF pass: one CAUSAL sequence per maximal parent-adjacent chain. A chain is a pure path + (only one child can start at its parent's end), so a chain token's causal prefix is exactly + its in-chain ancestors — correct for any DFS-preorder layout. + - FAT cross pass: each chain attends the maximal CONTIGUOUS prefix of its head's ancestor + path (walking down from the tree root while spans stay adjacent) — one fat-K sequence. + Chain-first layouts have fully contiguous ancestor paths, so this is their only cross pass. + - SKINNY cross passes: ancestors after the contiguity break (e.g. interior non-first children + in balanced trees) get one per-ancestor sequence each, grouped by break-index so every q row + appears at most once per pass (the merge-slot invariant).""" + ns = [int(x) for x in node_start] + nl = [int(x) for x in node_len] + par = [int(x) for x in node_parent] + N = len(ns) + total = max((ns[i] + nl[i] for i in range(N)), default=0) + + adj = [par[i] != -1 and ns[i] == ns[par[i]] + nl[par[i]] for i in range(N)] + chain_head = list(range(N)) + for i in range(N): + if adj[i]: + chain_head[i] = chain_head[par[i]] + + heads = sorted(set(chain_head)) + chain_end = {} + for i in range(N): + h = chain_head[i] + chain_end[h] = max(chain_end.get(h, 0), ns[i] + nl[i]) + + fat = [] # (q0, q1, a0, a1): contiguous ancestor-prefix sequences + skinny = {} # break-index -> list of (q0, q1, s0, s1) single-ancestor sequences + for h in heads: + if par[h] == -1: + continue + path, x = [], par[h] + while x != -1: + path.append(x) + x = par[x] + path.reverse() # tree root first + y0 = ns[path[0]] + y = y0 + nl[path[0]] + i = 1 + while i < len(path) and ns[path[i]] == y: + y += nl[path[i]] + i += 1 + q0, q1 = ns[h], chain_end[h] + fat.append((q0, q1, y0, y)) + for j, a in enumerate(path[i:]): + skinny.setdefault(j, []).append((q0, q1, ns[a], ns[a] + nl[a])) + + def _i32(x): + return torch.tensor(x, dtype=torch.int32, device=device) + + def _i64(x): + return torch.tensor(x, dtype=torch.long, device=device) + + passes = [] + # self pass: one causal sequence per CHAIN (q/k identity over [0, total)). + cu = [0] + for h in heads: + cu.append(cu[-1] + (chain_end[h] - ns[h])) + mx = max((chain_end[h] - ns[h] for h in heads), default=0) + passes.append((None, None, _i32(cu), _i32(cu), mx, mx, True)) + + def _build_cross(seqs): + seqs.sort(key=lambda c: c[0]) + if _SP_QSLICE: + # contiguous q view [qlo, qhi): real sequences + zero-K padding over the gaps. + qlo, qhi = seqs[0][0], seqs[-1][1] + kpos, cuq, cuk, gaps = [], [0], [0], [] + cur, mxq = qlo, 0 + for q0, q1, a0, a1 in seqs: + if q0 > cur: # gap: rows exist in the view but attend nothing + gaps.append((cur, q0)) + cuq.append(cuq[-1] + (q0 - cur)) + cuk.append(cuk[-1]) + mxq = max(mxq, q0 - cur) + kpos.append(torch.arange(a0, a1, dtype=torch.long, device=device)) + cuq.append(cuq[-1] + (q1 - q0)) + cuk.append(cuk[-1] + (a1 - a0)) + mxq = max(mxq, q1 - q0) + cur = q1 + mxk = max(c[3] - c[2] for c in seqs) + return (_QSlice(qlo, qhi, gaps), torch.cat(kpos), _i32(cuq), _i32(cuk), mxq, mxk, False) + qpos, kpos, cuq, cuk = [], [], [0], [0] + for q0, q1, a0, a1 in seqs: + qpos.append(torch.arange(q0, q1, dtype=torch.long, device=device)) + kpos.append(torch.arange(a0, a1, dtype=torch.long, device=device)) + cuq.append(cuq[-1] + (q1 - q0)) + cuk.append(cuk[-1] + (a1 - a0)) + mxq = max(c[1] - c[0] for c in seqs) + mxk = max(c[3] - c[2] for c in seqs) + return (torch.cat(qpos), torch.cat(kpos), _i32(cuq), _i32(cuk), mxq, mxk, False) + + if fat: + passes.append(_build_cross(fat)) + for j in sorted(skinny): + passes.append(_build_cross(skinny[j])) + return total, passes + + +def _forest_attention_plan(node_start, node_len, node_parent, device): + """Decompose a forest/tree into the flash passes the composed attention runs. + + Returns ``(total, passes)`` where ``passes`` is a list of + ``(q_idx, k_idx, cu_q, cu_k, max_q, max_k, causal)``: + * one SELF pass -- every node attends its own span causally (block-diagonal varlen over all + tokens), and + * one CROSS pass per ancestor depth level L -- tokens strictly below depth L attend their + level-L ancestor span, non-causally (DFS contiguity makes each ancestor's descendant tokens + a contiguous run). + A token attends ``{own node, causal} ∪ {each ancestor node, full}`` -- the union of the passes + it + appears in as a query. ``max_depth + 1`` passes total, independent of group count: depth-1 + forest + (stars) ⇒ self + 1 cross; arbitrary-depth trees ⇒ ``depth + 1``. ``node_*`` give the structure + (parents precede children; a subtree is a contiguous DFS run). + """ + ns = [int(x) for x in node_start] + nl = [int(x) for x in node_len] + par = [int(x) for x in node_parent] + N = len(ns) + total = max((ns[i] + nl[i] for i in range(N)), default=0) + + # PRECONDITION: node arrays must be DFS-preorder (each node's subtree is the contiguous run + # immediately after it). ``subtree_end`` below relies on this -- a non-DFS layout would silently + # attend the WRONG ancestor spans (branches lose interior-ancestor context -> corrupt logprobs). + # PackedTreeLayout only checks contiguity + parent depth[i]: + subtree_end[i] = ns[j] + nl[j] + j += 1 + + def _i32(x): + return torch.tensor(x, dtype=torch.int32, device=device) + + def _i64(x): + return torch.tensor(x, dtype=torch.long, device=device) + + passes = [] + # self pass: one block-diagonal causal varlen over every node's own span. Its q/k indices are + # the + # identity over [0, total), so they are left as ``None`` -- the kernel then uses q/k/v + # directly and + # skips a full-tensor gather/scatter every forward and backward (a real cost on big packed + # bins). + cu = [0] + for i in range(N): + cu.append(cu[-1] + nl[i]) + mnl = max(nl, default=0) + passes.append((None, None, _i32(cu), _i32(cu), mnl, mnl, True)) + + # cross passes: one per ancestor depth level. + for L in range(d_max): + qpos: List[int] = [] + kpos: List[int] = [] + cuq, cuk = [0], [0] + for a in range(N): + if depth[a] != L: + continue + qs, qe = (ns[a] + nl[a], subtree_end[a]) # strict descendants of a (contiguous, DFS) + if qe <= qs: + continue + qpos.extend(range(qs, qe)) + kpos.extend(range(ns[a], ns[a] + nl[a])) + cuq.append(cuq[-1] + (qe - qs)) + cuk.append(cuk[-1] + nl[a]) + if not qpos: + continue + mxq = max(cuq[i + 1] - cuq[i] for i in range(len(cuq) - 1)) + mxk = max(cuk[i + 1] - cuk[i] for i in range(len(cuk) - 1)) + passes.append((_i64(qpos), _i64(kpos), _i32(cuq), _i32(cuk), mxq, mxk, False)) + return total, passes + + +class _ComposedForestAttn(torch.autograd.Function): + """Composed forest/tree attention with an EXACT backward. + + The forward runs the ``_forest_attention_plan`` passes with flash and merges them by online + softmax (LSE) -- the union-softmax identity, so the forward is exact. The backward is the + delicate part: each pass's flash output is a sub-attention, and the *naive* + autograd-through-flash + drops the inter-pass normalizer-coupling term (flash exposes no gradient through its LSE), + giving + ~15% wrong q/k grads. We fix it by calling the low-level ``_flash_attn_varlen_backward`` per + pass + with ``dout = w_pass * do`` AND substituting the MERGED output ``o`` for the pass's own output: + flash uses ``out`` only to form the row-delta ``D = rowsum(dout ∘ out)``, which is exactly the + softmax ``G`` term, so feeding the merged ``o`` injects the global normalizer that was missing. + The result is the exact union-softmax score gradient ``P_ij (v_j·do − o·do)`` for every pass -- + dq, dk, dv all correct -- with no custom kernel (see TREE_PACKING_DESIGN.md §5).""" + + @staticmethod + def forward(ctx, q, k, v, node_start, node_len, node_parent, scale, full_context): + # q: [total, np, hn]; k, v: [total, ng, hn]; scale already resolved to a float. + """Evaluate independent causal and ancestor passes, then merge their softmax outputs.""" + from flash_attn import flash_attn_varlen_func + + total, passes, inv_maps, qidx64, slot_pass = _forest_attention_plan_cached( + node_start, node_len, node_parent, q.device, full_context=full_context + ) + np_, hn = q.shape[1], q.shape[2] + streams = _sp_streams() + outs, lses = [None] * len(passes), [None] * len(passes) + if streams is not None and len(passes) > 1: + # passes are independent until the merge: fan them out on side streams so the small + # cross passes hide under the big self pass. Their outputs are consumed back on the + # current stream after the join events (record_stream keeps the allocator honest). + cur = torch.cuda.current_stream() + ev_in = torch.cuda.Event() + ev_in.record(cur) + join = [] + for i, (q_idx, k_idx, cu_q, cu_k, mxq, mxk, causal) in enumerate(passes): + st = streams[i % len(streams)] + st.wait_event(ev_in) + with torch.cuda.stream(st): + qx = _sel_rows(q, q_idx) + if k_idx is None: + kx, vx = k, v + else: + kx, vx = _gather_kv(k, v, k_idx) + o, lse, _ = flash_attn_varlen_func( + qx, + kx, + vx, + cu_q, + cu_k, + mxq, + mxk, + softmax_scale=scale, + causal=causal, + return_attn_probs=True, + ) + o.record_stream(cur) + lse.record_stream(cur) + ev = torch.cuda.Event() + ev.record(st) + join.append(ev) + outs[i], lses[i] = o, lse + for ev in join: + cur.wait_event(ev) + else: + for i, (q_idx, k_idx, cu_q, cu_k, mxq, mxk, causal) in enumerate(passes): + qx = _sel_rows(q, q_idx) # None => identity (self pass) + if k_idx is None: + kx, vx = k, v + else: + kx, vx = _gather_kv(k, v, k_idx) + o, lse, _ = flash_attn_varlen_func( + qx, + kx, + vx, + cu_q, + cu_k, + mxq, + mxk, + softmax_scale=scale, + causal=causal, + return_attn_probs=True, + ) # o [Σq, np, hn], lse [np, Σq] + outs[i], lses[i] = o, lse + + if _SP_FUSED_MERGE and HAVE_TRITON and len(slot_pass) <= 7: + # single-kernel online-softmax merge: reads each pass output/LSE once, writes the + # merged output + final LSE once (fp32 in-register), replacing the eager fp32 + # upcast/mul/index_add round-trips below. + o_merged, lse_final = _merge_passes_triton( + slot_pass, inv_maps, outs, lses, total, np_, hn, q.dtype, q.device + ) + else: + # merged LSE per (head, token): logsumexp over every pass the token queries in. + lse_final = torch.full( + (np_, total), float("-inf"), device=q.device, dtype=torch.float32 + ) + for (q_idx, *_), lse in zip(passes, lses): + ls = lse.float() + if isinstance(q_idx, _QSlice): + # zero-K padding rows report LSE=+inf: neutralize before merging. + ls = torch.where(torch.isinf(ls), torch.full_like(ls, float("-inf")), ls) + lse_final[:, q_idx.lo : q_idx.hi] = torch.logaddexp( + lse_final[:, q_idx.lo : q_idx.hi], ls + ) + elif q_idx is None: + lse_final = torch.logaddexp(lse_final, ls) + else: + lse_final[:, q_idx] = torch.logaddexp(lse_final[:, q_idx], ls) + # merged output: sum_pass w_pass * o_pass, w_pass = exp(lse_pass - lse_final). + o_merged = torch.zeros(total, np_, hn, device=q.device, dtype=torch.float32) + for (q_idx, *_), o, lse in zip(passes, outs, lses): + ls = lse.float() + if isinstance(q_idx, _QSlice): + ls = torch.where(torch.isinf(ls), torch.full_like(ls, float("-inf")), ls) + lf = lse_final[:, q_idx.lo : q_idx.hi] + contrib = torch.exp(ls - lf).transpose(0, 1).unsqueeze(-1) * o.float() + o_merged[q_idx.lo : q_idx.hi] += contrib + continue + lf = lse_final if q_idx is None else lse_final.index_select(1, q_idx) + contrib = torch.exp(ls - lf).transpose(0, 1).unsqueeze(-1) * o.float() + if q_idx is None: + o_merged = o_merged + contrib + else: + o_merged.index_add_(0, q_idx, contrib) + o_merged = o_merged.to(q.dtype) + + ctx.save_for_backward(q, k, v, o_merged) + ctx.passes = passes + ctx.lses = lses + ctx.lse_final = lse_final + ctx.scale = scale + ctx.qidx64 = qidx64 + ctx.inv_maps = inv_maps + ctx.slot_pass = slot_pass + return o_merged + + @staticmethod + def backward(ctx, do): + """Accumulate pass gradients using the merged output for the global softmax correction.""" + from flash_attn.flash_attn_interface import _flash_attn_varlen_backward + + q, k, v, o_merged = ctx.saved_tensors + lse_final, scale = ctx.lse_final, ctx.scale + do = do.contiguous() + total, np_, hn = q.shape[0], q.shape[1], q.shape[2] + have_cached_row_maps = getattr(ctx, "qidx64", None) is not None + # Fused glue initializes KV gradients from an identity self pass. A full-context + # plan gathers repeated ancestor KV in its first pass and needs the general scatter. + use_backward_glue = ( + _sp_fused_backward_glue_effective() + and have_cached_row_maps + and ctx.passes[0][1] is None + ) + use_dq_assembly = _sp_fused_dq_assembly_effective(ctx.slot_pass) and have_cached_row_maps + + if use_backward_glue: + # Fused backward glue: (a) per-pass dout scaling w*do fused with the query gather in + # one kernel (the eager path materialized an fp32 exp/mul chain + an index_select per + # pass); (b) the self pass (always pass 0, identity indices over all tokens) INITIALIZES + # the fp32 accumulators instead of zeros+add; cross passes scatter-accumulate through a + # cast-fused kernel (no per-pass .float() temporaries). The flash calls are unchanged + # -- the exact-backward trick (merged o substituted for the pass output) is preserved. + ng = k.shape[1] + streams = _sp_streams() + cur = torch.cuda.current_stream() + results = [None] * len(ctx.passes) + join = [None] * len(ctx.passes) + ev_in = None + if streams is not None and len(ctx.passes) > 1: + ev_in = torch.cuda.Event() + ev_in.record(cur) + for i, ((q_idx, k_idx, cu_q, cu_k, mxq, mxk, causal), lse) in enumerate( + zip(ctx.passes, ctx.lses) + ): + st = None if ev_in is None else streams[i % len(streams)] + stream_ctx = torch.cuda.stream(st) if st is not None else nullcontext() + if st is not None: + st.wait_event(ev_in) + with stream_ctx: + qidx = ctx.qidx64[i] + rows = qidx.numel() + qx = _sel_rows(q, q_idx) + if k_idx is None: + kx, vx = k, v + else: + kx, vx = _gather_kv(k, v, k_idx) + ox = _sel_rows(o_merged, q_idx) + dox = torch.empty(rows, np_, hn, dtype=q.dtype, device=q.device) + _sp_scale_gather_kernel[((rows + 15) // 16, np_)]( + do, lse.contiguous(), lse_final, qidx, dox, rows, total, np_=np_, HN=hn + ) + dqx, dkx, dvx = ( + torch.empty_like(qx), + torch.empty_like(kx), + torch.empty_like(vx), + ) + _flash_attn_varlen_backward( + dox, + qx, + kx, + vx, + ox, + lse, + dqx, + dkx, + dvx, + cu_q, + cu_k, + mxq, + mxk, + 0.0, + scale, + causal, + -1, + -1, + 0.0, + None, + _SP_DETERMINISTIC_BACKWARD, + None, + False, + ) + results[i] = (dqx, dkx, dvx, qidx, k_idx, rows, q_idx) + if st is not None: + dqx.record_stream(cur) + dkx.record_stream(cur) + dvx.record_stream(cur) + ev = torch.cuda.Event() + ev.record(st) + join[i] = ev + # k/v accumulate in pass order on the current stream (scatters are read-modify-write + # on shared fp32 accumulators, ng=8 so the buffers are small); dq is assembled by one + # gather-side merge kernel over all pass contributions (opt10) once every pass joins. + dk = dv = None + dqxs = [None] * len(ctx.passes) + for i, res in enumerate(results): + if join[i] is not None: + cur.wait_event(join[i]) + dqx, dkx, dvx, qidx, k_idx, rows, q_idx = res + dqxs[i] = dqx + if i == 0: + # self pass: identity over all tokens -> direct init, no zeros/scatter. + dk = dkx.float() + dv = dvx.float() + else: + kro = k_idx.numel() + _sp_scatter_accum_kernel[((kro + 15) // 16, ng)]( + dk, dkx, k_idx, kro, n_=ng, HN=hn + ) + _sp_scatter_accum_kernel[((kro + 15) // 16, ng)]( + dv, dvx, k_idx, kro, n_=ng, HN=hn + ) + if use_dq_assembly: + dq = _merge_dq_triton( + ctx.slot_pass, ctx.inv_maps, dqxs, total, np_, hn, q.dtype, q.device + ) + else: + dq = torch.zeros(q.shape, device=q.device, dtype=torch.float32) + for dqx, result in zip(dqxs, results): + q_idx = result[-1] + if q_idx is None: + dq += dqx.float() + elif isinstance(q_idx, _QSlice): + dq[q_idx.lo : q_idx.hi] += dqx.float() + else: + dq.index_add_(0, q_idx, dqx.float()) + dq = dq.to(q.dtype) + return dq, dk.to(k.dtype), dv.to(v.dtype), None, None, None, None, None + + dq = None if use_dq_assembly else torch.zeros(q.shape, device=q.device, dtype=torch.float32) + dqxs = [] if use_dq_assembly else None + dk = torch.zeros(k.shape, device=k.device, dtype=torch.float32) + dv = torch.zeros(v.shape, device=v.device, dtype=torch.float32) + for (q_idx, k_idx, cu_q, cu_k, mxq, mxk, causal), lse in zip(ctx.passes, ctx.lses): + qx = _sel_rows(q, q_idx) # q_idx None => identity (self pass) + if k_idx is None: + kx, vx = k, v + else: + kx, vx = _gather_kv(k, v, k_idx) + ox = _sel_rows(o_merged, q_idx) # MERGED output -> exact + if isinstance(q_idx, _QSlice): + lf = lse_final[:, q_idx.lo : q_idx.hi] + dox_full = do[q_idx.lo : q_idx.hi] + ls = lse.float() + ls = torch.where(torch.isinf(ls), torch.full_like(ls, float("-inf")), ls) + w = torch.exp(ls - lf) # 0 on zero-K padding rows + else: + lf = lse_final if q_idx is None else lse_final.index_select(1, q_idx) + dox_full = do if q_idx is None else do.index_select(0, q_idx) + w = torch.exp(lse.float() - lf) + dox = (w.transpose(0, 1).unsqueeze(-1) * dox_full).to(q.dtype) + dqx, dkx, dvx = (torch.empty_like(qx), torch.empty_like(kx), torch.empty_like(vx)) + _flash_attn_varlen_backward( + dox, + qx, + kx, + vx, + ox, + lse, + dqx, + dkx, + dvx, + cu_q, + cu_k, + mxq, + mxk, + 0.0, + scale, + causal, + -1, + -1, + 0.0, + None, + _SP_DETERMINISTIC_BACKWARD, + None, + False, + ) + if use_dq_assembly: + dqxs.append(dqx) + else: + if q_idx is None: + dq += dqx.float() + elif isinstance(q_idx, _QSlice): + dq[q_idx.lo : q_idx.hi] += dqx.float() + else: + dq.index_add_(0, q_idx, dqx.float()) + if k_idx is None: + dk += dkx.float() + dv += dvx.float() + else: + dk.index_add_(0, k_idx, dkx.float()) + dv.index_add_(0, k_idx, dvx.float()) + if use_dq_assembly: + dq = _merge_dq_triton( + ctx.slot_pass, ctx.inv_maps, dqxs, total, np_, hn, q.dtype, q.device + ) + else: + dq = dq.to(q.dtype) + return dq, dk.to(k.dtype), dv.to(v.dtype), None, None, None, None, None + + +def flash_composed_forest_attention_fused( + query, key, value, node_start, node_len, node_parent, scale=None, *, full_context: bool = False +): + """Level-decomposed forest/tree attention, fused to ``max_depth + 1`` flash passes per bin. + + A token attends ``{its own node, causally} ∪ {each ancestor node, fully}``. Decomposed by + ancestor + depth level (see :func:`_forest_attention_plan`) and merged by online softmax, so cost is + independent of group count (depth-1 stars ⇒ 1 self + 1 cross pass) -- vs the per-group loop's + ``~3 * #groups`` launches. Forward AND backward are exact via :class:`_ComposedForestAttn`. + + ``query`` is ``[sq, b, np, hn]`` and ``key``/``value`` ``[sq, b, ng, hn]`` (b == 1). Trailing + pad + positions (beyond the last node) get zero outputs and are never read downstream. Returns + ``[sq, b, np*hn]``; same default scale (1/sqrt(hn)) as the flex/loop paths. + """ + sq, b, np_, hn = query.shape + assert b == 1, "shared-prefix packing uses a single packed sequence (b == 1)" + total = max((int(s) + int(l) for s, l in zip(node_start, node_len)), default=0) + scale = scale if scale is not None else hn**-0.5 + out = _ComposedForestAttn.apply( + query[:total, 0], + key[:total, 0], + value[:total, 0], + node_start, + node_len, + node_parent, + scale, + full_context, + ) # [total, np, hn] + if sq > total: + out = torch.cat([out, out.new_zeros(sq - total, np_, hn)], dim=0) + return out.reshape(sq, 1, np_ * hn).contiguous() + + +def _forest_to_nodes(forest): + """Expand a depth-1 ``forest`` list into flat node arrays (PackedTreeLayout structure). + + ``forest`` is ``[(token_offset, prefix_len, completion_lens), ...]``. Each group becomes a root + node (the prefix span, parent -1) followed by one child node per completion. Node order is DFS + (root then its children), as the fused kernel requires. + """ + node_start, node_len, node_parent = [], [], [] + for off, prefix_len, completion_lens in forest: + root = len(node_start) + node_start.append(int(off)) + node_len.append(int(prefix_len)) + node_parent.append(-1) + pos = int(off) + int(prefix_len) + for c in completion_lens: + node_start.append(pos) + node_len.append(int(c)) + node_parent.append(root) + pos += int(c) + return node_start, node_len, node_parent + + +def flash_composed_forest_attention( + query, key, value, forest, scale=None, *, full_context: bool = False +): + """Forest (multi-group / Case-1) shared-prefix attention -- fused, level-decomposed. + + ``forest`` is a list of ``(token_offset, prefix_len, completion_lens)``, one per group packed + into this bin. Expanded to flat node arrays and dispatched to + :func:`flash_composed_forest_attention_fused`, which fuses ALL groups into ``max_depth + 1`` + flash calls (depth-1 forest -> 1 self + 1 cross + merge) regardless of group count -- vs the + old per-group loop's ``~3 * #groups`` launches, which scaled badly with many small groups. + See :func:`flash_composed_forest_attention_loop` for the reference per-group form (kept for + parity testing). + """ + node_start, node_len, node_parent = _forest_to_nodes(forest) + return flash_composed_forest_attention_fused( + query, key, value, node_start, node_len, node_parent, scale=scale, full_context=full_context + ) + + +def _undo_cp_zigzag(input_: torch.Tensor, cp_size: int) -> torch.Tensor: + """Convert rank-major CP zigzag chunks into canonical global token order. + + Standard context parallelism gives rank ``r`` chunks ``(r, 2*C-r-1)``. An + all-to-all from sequence to head parallelism concatenates those rank-local pairs in + rank order, i.e. ``0,last,1,last-1,...``. Mamba uses the same permutation before + its sequential scan; forest attention needs it for its global node offsets too. + """ + if cp_size < 1 or input_.shape[0] % (2 * cp_size): + raise ValueError("CP zigzag sequence length must be divisible by 2 * CP size") + chunks = torch.chunk(input_, chunks=2 * cp_size, dim=0) + order = [2 * index for index in range(cp_size)] + [ + 2 * cp_size - 2 * index - 1 for index in range(cp_size) + ] + return torch.cat([chunks[index] for index in order], dim=0) + + +def _redo_cp_zigzag(input_: torch.Tensor, cp_size: int) -> torch.Tensor: + """Convert canonical global token order back to rank-major CP zigzag chunks.""" + if cp_size < 1 or input_.shape[0] % (2 * cp_size): + raise ValueError("CP zigzag sequence length must be divisible by 2 * CP size") + chunks = torch.chunk(input_, chunks=2 * cp_size, dim=0) + order = [None] * (2 * cp_size) + order[::2] = range(cp_size) + order[1::2] = reversed(range(cp_size, 2 * cp_size)) + return torch.cat([chunks[index] for index in order], dim=0) + + +def _cp_local_kv_head_slice(query_heads: int, kv_heads: int, cp_size: int, cp_rank: int) -> slice: + """Return the global KV-head span serving one CP rank's contiguous Q-head span. + + The fused flash primitive supports GQA, but its local Q-to-KV grouping must exactly + match the global grouping after Q heads are sharded over CP. Some valid layouts + (including 32 Q heads / 2 KV heads / CP4) replicate a single KV head on multiple CP + ranks. Layouts whose CP boundary cuts unequal pieces from several KV groups fail + closed rather than silently changing attention semantics. + """ + if min(query_heads, kv_heads, cp_size) < 1 or not 0 <= cp_rank < cp_size: + raise ValueError( + f"invalid shared-prefix CP head geometry: {query_heads=}, {kv_heads=}, " + f"{cp_size=}, {cp_rank=}" + ) + if query_heads % cp_size: + raise NotImplementedError( + f"shared-prefix CP attention requires {query_heads=} divisible by {cp_size=}" + ) + if query_heads % kv_heads: + raise NotImplementedError( + "shared-prefix CP attention requires an integral global Q/KV head ratio" + ) + + local_query_heads = query_heads // cp_size + queries_per_kv = query_heads // kv_heads + query_start = cp_rank * local_query_heads + global_mapping = [ + query_index // queries_per_kv + for query_index in range(query_start, query_start + local_query_heads) + ] + kv_start = global_mapping[0] + kv_stop = global_mapping[-1] + 1 + local_kv_heads = kv_stop - kv_start + if local_query_heads % local_kv_heads: + raise NotImplementedError( + "shared-prefix CP attention cannot express this Q/KV grouping with local flash GQA" + ) + local_queries_per_kv = local_query_heads // local_kv_heads + expected_mapping = [ + kv_start + local_index // local_queries_per_kv for local_index in range(local_query_heads) + ] + if global_mapping != expected_mapping: + raise NotImplementedError( + "shared-prefix CP attention boundary cuts unequal portions of multiple KV groups" + ) + return slice(kv_start, kv_stop) + + +def _cp_kv_head_slices_for_destinations( + query_heads: int, kv_heads: int, cp_size: int +) -> tuple[slice, ...]: + """Return equal-width KV-head blocks in all-to-all destination order. + + Each destination owns one contiguous Q-head block. Its KV block can overlap + another destination's block for GQA, but the equal-split all-to-all requires + every destination block to contain the same number of heads. + """ + destination_slices = tuple( + _cp_local_kv_head_slice(query_heads, kv_heads, cp_size, cp_rank) + for cp_rank in range(cp_size) + ) + destination_widths = tuple( + head_slice.stop - head_slice.start for head_slice in destination_slices + ) + if len(set(destination_widths)) != 1: + raise NotImplementedError( + "shared-prefix CP attention requires equal KV-head widths for every " + f"all-to-all destination; got {destination_widths}" + ) + return destination_slices + + +def _cp_pack_destination_head_slices( + input_: torch.Tensor, destination_slices: tuple[slice, ...] +) -> torch.Tensor: + """Pack destination-specific head blocks before an equal-split all-to-all. + + Overlapping slices are intentionally concatenated more than once. Autograd + accumulates their backward contributions into the shared source heads, which + preserves GQA KV-gradient multiplicity without materializing every KV head on + every destination. + """ + if not destination_slices: + raise ValueError("shared-prefix CP attention requires at least one destination") + head_count = input_.shape[2] + widths = [] + for head_slice in destination_slices: + if ( + head_slice.step not in (None, 1) + or head_slice.start is None + or head_slice.stop is None + or not 0 <= head_slice.start < head_slice.stop <= head_count + ): + raise ValueError( + "shared-prefix CP attention destination slices must be non-empty, " + "unit-stride, and within the KV-head dimension" + ) + widths.append(head_slice.stop - head_slice.start) + if len(set(widths)) != 1: + raise NotImplementedError( + "shared-prefix CP attention requires equal KV-head widths for every " + f"all-to-all destination; got {tuple(widths)}" + ) + return torch.cat([input_[:, :, head_slice, :] for head_slice in destination_slices], dim=2) + + +def _cp_sequence_to_head_parallel( + input_: torch.Tensor, + cp_group: torch.distributed.ProcessGroup, + *, + destination_head_slices: tuple[slice, ...] | None = None, +) -> torch.Tensor: + """Change ``[S/C,1,H,D]`` zigzag sequence shards to canonical head shards.""" + if input_.ndim != 4 or input_.shape[1] != 1: + raise ValueError("shared-prefix CP attention requires Q/K/V shape [S/C,1,H,D]") + cp_size = cp_group.size() + if input_.shape[0] % 2: + raise ValueError( + "shared-prefix CP attention requires an even local sequence length for zigzag CP" + ) + if destination_head_slices is not None: + if len(destination_head_slices) != cp_size: + raise ValueError( + "shared-prefix CP attention requires one head slice per all-to-all destination" + ) + input_ = _cp_pack_destination_head_slices(input_, destination_head_slices) + elif input_.shape[2] % cp_size: + raise NotImplementedError( + "shared-prefix CP attention requires Q heads divisible by context parallel size" + ) + + local_sequence, batch, heads, head_dim = input_.shape + exchanged = all_to_all_sp2hp( + input_.reshape(local_sequence, batch, heads * head_dim), group=cp_group + ) + global_sequence = local_sequence * cp_size + local_heads = heads // cp_size + exchanged = exchanged.reshape(global_sequence, batch, local_heads, head_dim) + return _undo_cp_zigzag(exchanged, cp_size) + + +def _cp_head_to_sequence_parallel( + input_: torch.Tensor, cp_group: torch.distributed.ProcessGroup +) -> torch.Tensor: + """Invert :func:`_cp_sequence_to_head_parallel` for attention output.""" + if input_.ndim != 4 or input_.shape[1] != 1: + raise ValueError("shared-prefix CP attention output must have shape [S,1,H/C,D]") + cp_size = cp_group.size() + global_sequence, batch, local_heads, head_dim = input_.shape + if global_sequence % (2 * cp_size): + raise ValueError("shared-prefix CP global sequence length must be divisible by 2 * CP size") + rank_major = _redo_cp_zigzag(input_, cp_size) + exchanged = all_to_all_hp2sp( + rank_major.reshape(global_sequence, batch, local_heads * head_dim), group=cp_group + ) + return exchanged.reshape(global_sequence // cp_size, batch, local_heads * cp_size, head_dim) + + +def flash_composed_forest_attention_cp( + query, + key, + value, + forest, + *, + cp_group: torch.distributed.ProcessGroup, + scale=None, + full_context: bool = False, +): + """Exact forest attention for standard zigzag context-parallel sequence shards. + + Q heads are sharded across CP while each destination receives only its required + K/V heads. Overlapping GQA slices remain differentiably replicated. All tensors + are converted to canonical global sequence order, then the existing optimized + exact-forward/exact-backward forest primitive runs on the local head shard. Its + output is transformed back to the caller's CP-local zigzag token order. No global + logits, hidden states, or full-KV activation bases are materialized. + """ + cp_size = cp_group.size() + if cp_size == 1: + return flash_composed_forest_attention( + query, key, value, forest, scale=scale, full_context=full_context + ) + if any(tensor.ndim != 4 for tensor in (query, key, value)): + raise ValueError("shared-prefix CP attention requires Q/K/V rank-4 tensors") + if any(tensor.shape[1] != 1 for tensor in (query, key, value)): + raise ValueError("shared-prefix CP attention requires Q/K/V batch size 1") + if not query.shape[0] == key.shape[0] == value.shape[0]: + raise ValueError("shared-prefix CP attention requires aligned local Q/K/V sequences") + if key.shape[2] != value.shape[2]: + raise ValueError("shared-prefix CP attention requires equal K/V head counts") + if not query.shape[3] == key.shape[3] == value.shape[3]: + raise ValueError("shared-prefix CP attention requires equal Q/K/V head dimensions") + if not query.dtype == key.dtype == value.dtype: + raise ValueError("shared-prefix CP attention requires equal Q/K/V dtypes") + if not query.device == key.device == value.device: + raise ValueError("shared-prefix CP attention requires Q/K/V on the same device") + + query_heads = query.shape[2] + kv_heads = key.shape[2] + destination_kv_slices = _cp_kv_head_slices_for_destinations(query_heads, kv_heads, cp_size) + query_global = _cp_sequence_to_head_parallel(query, cp_group) + key_global = _cp_sequence_to_head_parallel( + key, cp_group, destination_head_slices=destination_kv_slices + ) + value_global = _cp_sequence_to_head_parallel( + value, cp_group, destination_head_slices=destination_kv_slices + ) + + output = flash_composed_forest_attention( + query_global, key_global, value_global, forest, scale=scale, full_context=full_context + ) + output = output.reshape(query_global.shape[0], 1, query_global.shape[2], query_global.shape[3]) + output = _cp_head_to_sequence_parallel(output, cp_group) + return output.reshape(output.shape[0], 1, -1).contiguous() diff --git a/megatron/core/models/hybrid/shared_prefix_layout.py b/megatron/core/models/hybrid/shared_prefix_layout.py new file mode 100644 index 00000000000..69ced9d8fbe --- /dev/null +++ b/megatron/core/models/hybrid/shared_prefix_layout.py @@ -0,0 +1,311 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +"""Token layouts for independent shared-prefix hybrid execution groups.""" + +from __future__ import annotations + +from collections.abc import Iterator, Sequence +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +from torch import Tensor + + +@dataclass(frozen=True) +class SharedPrefixLayout: + """One exact-prompt star packed as ``[prefix, completion_1, ..., completion_G]``.""" + + prefix_len: int + completion_lens: Sequence[int] + logical_completion_lens: Sequence[int] | None = None + padding_multiple: int | None = None + + def __post_init__(self) -> None: + prefix_len = int(self.prefix_len) + completion_lens = tuple(int(length) for length in self.completion_lens) + if prefix_len < 1: + raise ValueError("shared-prefix layout requires a non-empty prefix") + if not completion_lens or any(length < 1 for length in completion_lens): + raise ValueError("shared-prefix layout requires one or more non-empty completions") + object.__setattr__(self, "prefix_len", prefix_len) + object.__setattr__(self, "completion_lens", completion_lens) + if self.logical_completion_lens is not None: + logical_completion_lens = tuple(int(length) for length in self.logical_completion_lens) + if len(logical_completion_lens) != len(completion_lens) or any( + logical < 1 or logical > physical + for logical, physical in zip(logical_completion_lens, completion_lens, strict=True) + ): + raise ValueError( + "logical completion lengths must be positive, match the physical " + "branch count, and not exceed physical completion lengths" + ) + object.__setattr__(self, "logical_completion_lens", logical_completion_lens) + if (self.logical_completion_lens is None) != (self.padding_multiple is None): + raise ValueError( + "logical completion lengths and padding_multiple must be provided together" + ) + if self.padding_multiple is not None: + if isinstance(self.padding_multiple, bool) or not isinstance( + self.padding_multiple, int + ): + raise ValueError("shared-prefix padding_multiple must be an integer") + if self.padding_multiple < 1: + raise ValueError("shared-prefix padding_multiple must be positive") + + @property + def total_len(self) -> int: + """Return the physical token count of this star.""" + return self.prefix_len + sum(self.completion_lens) + + def iter_roots(self) -> Iterator[tuple[int, SharedPrefixLayout]]: + """Yield canonical offsets and independent prompt stars.""" + yield 0, self + + @property + def dense_branch_lengths(self) -> tuple[int, ...]: + """Return each completion length including its independent prompt copy.""" + return tuple(self.prefix_len + length for length in self.completion_lens) + + @property + def forest(self) -> list[tuple[int, int, list[int]]]: + """Represent this star as a root and its completion lengths.""" + return [(0, self.prefix_len, list(self.completion_lens))] + + def completion_slices(self) -> tuple[slice, ...]: + """Return physical slices for the disjoint completion segments.""" + slices = [] + start = self.prefix_len + for length in self.completion_lens: + slices.append(slice(start, start + length)) + start += length + return tuple(slices) + + def dense_branch_indices(self, device: torch.device | str) -> tuple[Tensor, ...]: + """Return global-star indices for conventional prompt-completion sequences. + + MTP's token shifts are defined on an ordinary causal sequence and must + never cross between sibling completion branches. These indices provide + the exact inverse of prompt deduplication: the prompt is repeated once + for each physical completion, including that completion's ordinary + per-sequence padding. + """ + prompt = torch.arange(self.prefix_len, device=device, dtype=torch.long) + return tuple( + torch.cat( + (prompt, torch.arange(branch.start, branch.stop, device=device, dtype=torch.long)) + ) + for branch in self.completion_slices() + ) + + def position_ids(self, device: torch.device | str) -> Tensor: + """Prefix-continued RoPE positions for the packed star.""" + pieces = [torch.arange(self.prefix_len, device=device, dtype=torch.long)] + pieces.extend( + torch.arange(self.prefix_len, self.prefix_len + length, device=device, dtype=torch.long) + for length in self.completion_lens + ) + return torch.cat(pieces) + + def padded_position_ids(self, physical_len: int, device: torch.device | str) -> Tensor: + """Return global star positions plus inert positions for trailing CP padding.""" + physical_len = int(physical_len) + if physical_len < self.total_len: + raise ValueError( + f"physical length {physical_len} is shorter than layout {self.total_len}" + ) + positions = self.position_ids(device) + if physical_len == self.total_len: + return positions + return torch.cat( + [positions, torch.zeros(physical_len - self.total_len, device=device, dtype=torch.long)] + ) + + def padded_token_multiplicities( + self, + physical_len: int, + device: torch.device | str, + *, + exclude_sequence_padding: bool = False, + ) -> Tensor: + """Dense-baseline multiplicity for each physical shared-prefix token. + + Prompt tokens occur once per completion in a conventional rollout batch. + Branch tokens, including ordinary per-sequence padding, have unit + multiplicity unless exclude_sequence_padding is requested. Trailing topology-only padding + is inert. + """ + physical_len = int(physical_len) + if physical_len < self.total_len: + raise ValueError( + f"physical length {physical_len} is shorter than layout {self.total_len}" + ) + multiplicities = torch.cat( + [ + torch.full( + (self.prefix_len,), + len(self.completion_lens), + device=device, + dtype=torch.float32, + ), + torch.ones(sum(self.completion_lens), device=device, dtype=torch.float32), + ] + ) + if exclude_sequence_padding: + # Hybridep's dense input mask excludes ordinary sequence padding. + # Keep the legacy unmasked convention as the default for other callers. + logical_lengths = self.logical_completion_lens or self.completion_lens + offset = self.prefix_len + for logical, physical in zip(logical_lengths, self.completion_lens): + if logical < physical: + multiplicities[offset + logical : offset + physical] = 0 + offset += physical + if physical_len == self.total_len: + return multiplicities + return torch.cat( + [ + multiplicities, + torch.zeros(physical_len - self.total_len, device=device, dtype=torch.float32), + ] + ) + + @staticmethod + def cp_local_indices( + physical_len: int, cp_size: int, cp_rank: int, device: torch.device | str + ) -> Tensor: + """Global indices owned by one rank under standard two-chunk CP zigzag.""" + physical_len, cp_size, cp_rank = int(physical_len), int(cp_size), int(cp_rank) + if cp_size < 1 or not 0 <= cp_rank < cp_size: + raise ValueError(f"invalid CP geometry: {cp_size=}, {cp_rank=}") + if physical_len % (2 * cp_size): + raise ValueError( + f"physical length {physical_len} must be divisible by 2 * CP size {cp_size}" + ) + chunk = physical_len // (2 * cp_size) + front = torch.arange(cp_rank * chunk, (cp_rank + 1) * chunk, device=device) + back_chunk = 2 * cp_size - cp_rank - 1 + back = torch.arange(back_chunk * chunk, (back_chunk + 1) * chunk, device=device) + return torch.cat([front, back]).to(torch.long) + + +@dataclass(frozen=True) +class SharedPrefixForestLayout: + """Independent prompt stars concatenated into one hybrid forward. + + Each branch retains its ordinary dense padding. Only the complete forest + receives topology padding; roots need no individual alignment because the + Mamba CP transform restores canonical token order before scanning them. + + ``mtp_loss_group_root_counts`` partitions consecutive roots into the original + independently normalized MTP forwards. Empty counts preserve one group per + root. Grouping loss normalization never joins causal sequences. + """ + + roots: tuple[SharedPrefixLayout, ...] + mtp_loss_group_root_counts: tuple[int, ...] = () + + def __post_init__(self) -> None: + object.__setattr__(self, "roots", tuple(self.roots)) + if not self.roots or any(not isinstance(root, SharedPrefixLayout) for root in self.roots): + raise ValueError("a shared-prefix forest requires one or more star layouts") + if len({root.padding_multiple for root in self.roots}) != 1: + raise ValueError("all shared-prefix roots must use the same padding multiple") + if isinstance(self.mtp_loss_group_root_counts, (str, bytes)): + raise ValueError("MTP loss group root counts must be a sequence of integers") + try: + counts = tuple(self.mtp_loss_group_root_counts) + except TypeError: + raise ValueError("MTP loss group root counts must be a sequence of integers") from None + if any( + isinstance(count, bool) or not isinstance(count, int) or count < 1 for count in counts + ): + raise ValueError("MTP loss group root counts must be positive integers") + if counts and sum(counts) != len(self.roots): + raise ValueError("MTP loss group root counts must partition every forest root") + object.__setattr__(self, "mtp_loss_group_root_counts", counts) + + @property + def total_len(self) -> int: + """Return the physical token count across all independent stars.""" + return sum(root.total_len for root in self.roots) + + @property + def padding_multiple(self) -> int | None: + """Return the shared branch-padding multiple, if specified.""" + return self.roots[0].padding_multiple + + @property + def dense_branch_lengths(self) -> tuple[int, ...]: + """Return dense branch lengths in forest order.""" + return tuple(length for root in self.roots for length in root.dense_branch_lengths) + + @property + def mtp_loss_group_lengths(self) -> tuple[int, ...]: + """Global dense-expanded token lengths of consecutive MTP loss groups. + + Ordinary per-branch padding is included; forest-only topology padding + never enters MTP. Each length is divided by CP size after branch packing. + """ + counts = self.mtp_loss_group_root_counts or (1,) * len(self.roots) + lengths = [] + start = 0 + for count in counts: + lengths.append( + sum(sum(root.dense_branch_lengths) for root in self.roots[start : start + count]) + ) + start += count + return tuple(lengths) + + def iter_roots(self) -> Iterator[tuple[int, SharedPrefixLayout]]: + """Yield each root with its global token offset.""" + offset = 0 + for root in self.roots: + yield offset, root + offset += root.total_len + + @property + def forest(self) -> list[tuple[int, int, list[int]]]: + """Return root offsets and lengths for the attention forest.""" + return [ + (offset, root.prefix_len, list(root.completion_lens)) + for offset, root in self.iter_roots() + ] + + def dense_branch_indices(self, device: torch.device | str) -> tuple[Tensor, ...]: + """Return independent dense branches without crossing root boundaries.""" + return tuple( + indices + offset + for offset, root in self.iter_roots() + for indices in root.dense_branch_indices(device) + ) + + def position_ids(self, device: torch.device | str) -> Tensor: + """Return per-branch positions in physical forest order.""" + return torch.cat([root.position_ids(device) for root in self.roots]) + + def padded_position_ids(self, physical_len: int, device: torch.device | str) -> Tensor: + """Extend forest positions to the requested physical length.""" + if physical_len < self.total_len: + raise ValueError("physical length is shorter than shared-prefix forest") + return F.pad(self.position_ids(device), (0, physical_len - self.total_len)) + + def padded_token_multiplicities( + self, + physical_len: int, + device: torch.device | str, + *, + exclude_sequence_padding: bool = False, + ) -> Tensor: + """Count logical token copies represented by each physical forest row.""" + if physical_len < self.total_len: + raise ValueError("physical length is shorter than shared-prefix forest") + weights = torch.cat( + [ + root.padded_token_multiplicities( + root.total_len, device, exclude_sequence_padding=exclude_sequence_padding + ) + for root in self.roots + ] + ) + return F.pad(weights, (0, physical_len - self.total_len)) + + cp_local_indices = staticmethod(SharedPrefixLayout.cp_local_indices) diff --git a/megatron/core/ssm/mamba_branch_layout.py b/megatron/core/ssm/mamba_branch_layout.py new file mode 100644 index 00000000000..0d26dd199de --- /dev/null +++ b/megatron/core/ssm/mamba_branch_layout.py @@ -0,0 +1,113 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Branch layout copies with one destination allocation in each backward. + +Ordinary differentiable slice assignments build a CopySlices chain. Its backward +can repeatedly materialize the entire branch rectangle or packed input. These +operators keep the same forward copies and explicitly assemble each input +gradient once. They do not change convolution or recurrent scan arithmetic. +""" + +import torch +from torch import Tensor + + +class _PackBranches(torch.autograd.Function): + @staticmethod + def forward(ctx, projected, prefix_len, tail_len, lengths): + """Pack the shared prompt tail and completions into a padded branch rectangle.""" + if projected.ndim != 3 or projected.shape[1] != 1: + raise ValueError("Mamba branch packing requires [sequence, 1, channels]") + if not lengths or min(lengths) < 0 or not 0 <= tail_len <= prefix_len: + raise ValueError("Invalid Mamba prefix tail or completion lengths") + if prefix_len + sum(lengths) != projected.shape[0]: + raise ValueError("Mamba branches must cover the physical projected input") + ctx.prefix_len = prefix_len + ctx.tail_len = tail_len + ctx.lengths = lengths + ctx.input_shape = projected.shape + branches = projected.new_zeros(tail_len + max(lengths), len(lengths), projected.shape[-1]) + start = prefix_len + for branch, length in enumerate(lengths): + if tail_len: + branches[:tail_len, branch] = projected[prefix_len - tail_len : prefix_len, 0] + branches[tail_len : tail_len + length, branch] = projected[start : start + length, 0] + start += length + return branches + + @staticmethod + def backward(ctx, grad_branches): + """Restore completion gradients and sum shared tail contributions in reverse order.""" + gradient = grad_branches.new_zeros(ctx.input_shape) + end = ctx.input_shape[0] + # Reverse branch order follows the original slice-assignment graph. Keep + # accumulation in the input dtype, including BF16 rounding at each add. + for branch in range(len(ctx.lengths) - 1, -1, -1): + length = ctx.lengths[branch] + start = end - length + gradient[start:end, 0] = grad_branches[ctx.tail_len : ctx.tail_len + length, branch] + if ctx.tail_len: + gradient[ctx.prefix_len - ctx.tail_len : ctx.prefix_len, 0].add_( + grad_branches[: ctx.tail_len, branch] + ) + end = start + return gradient, None, None, None + + +class _MergeBranches(torch.autograd.Function): + @staticmethod + def forward(ctx, prefix_head, branches, tail_len, lengths): + """Join the unique prompt and unpadded completion rows in physical order.""" + if prefix_head.ndim != 3 or prefix_head.shape[1] != 1: + raise ValueError("Mamba prefix head requires [sequence, 1, channels]") + if ( + branches.ndim != 3 + or branches.shape[1] != len(lengths) + or branches.shape[-1] != prefix_head.shape[-1] + or not lengths + or min(lengths) < 0 + or tail_len < 0 + or tail_len + max(lengths) != branches.shape[0] + ): + raise ValueError("Mamba branch rectangle does not match its layout") + ctx.head_len = prefix_head.shape[0] + ctx.branch_shape = branches.shape + ctx.tail_len = tail_len + ctx.lengths = lengths + return torch.cat( + [prefix_head, branches[:tail_len, :1]] + + [ + branches[tail_len : tail_len + length, branch : branch + 1] + for branch, length in enumerate(lengths) + ], + dim=0, + ) + + @staticmethod + def backward(ctx, gradient): + """Distribute gradients to the unique prompt head and selected branch rows.""" + grad_head = gradient[: ctx.head_len] + grad_branches = gradient.new_zeros(ctx.branch_shape) + start = ctx.head_len + grad_branches[: ctx.tail_len, :1] = gradient[start : start + ctx.tail_len] + start += ctx.tail_len + for branch, length in enumerate(ctx.lengths): + grad_branches[ctx.tail_len : ctx.tail_len + length, branch : branch + 1] = gradient[ + start : start + length + ] + start += length + return grad_head, grad_branches, None, None + + +def pack_mamba_branches( + projected: Tensor, *, prefix_len: int, tail_len: int, completion_lens: tuple[int, ...] +) -> Tensor: + """Copy a shared tail and disjoint completions into a padded branch batch.""" + return _PackBranches.apply(projected, prefix_len, tail_len, completion_lens) + + +def merge_mamba_branches( + prefix_head: Tensor, branches: Tensor, *, tail_len: int, completion_lens: tuple[int, ...] +) -> Tensor: + """Restore one prefix and all completions, ignoring sibling tail duplicates.""" + return _MergeBranches.apply(prefix_head, branches, tail_len, completion_lens) diff --git a/megatron/core/ssm/mamba_context_parallel.py b/megatron/core/ssm/mamba_context_parallel.py index 97259dcc0c4..ea8f3098405 100644 --- a/megatron/core/ssm/mamba_context_parallel.py +++ b/megatron/core/ssm/mamba_context_parallel.py @@ -141,11 +141,19 @@ def __init__( # either 1 or `ngroups_local_tp // cp_size` def pre_conv_ssm( - self, input_: torch.Tensor, packed_seq_params: Optional[PackedSeqParams] = None + self, + input_: torch.Tensor, + packed_seq_params: Optional[PackedSeqParams] = None, + *, + include_gate: bool = True, ) -> torch.Tensor: - """Method to be applied before the convolution and SSM""" + """Prepare convolution/SSM fields, optionally leaving the gate local. + + With ``include_gate=False``, return only canonical x/B/C/dt. The caller + retains the original TP-local gate and applies it after post_conv_ssm. + """ if self.cp_size == 1: - return input_ + return input_ if include_gate else input_[..., self.d_inner_local_tp :] z, x, B, C, dt = torch.split( input_, @@ -162,7 +170,8 @@ def pre_conv_ssm( # TODO (duncan): Can the some or all of the all_to_alls be combined? # [l_global//cp, b, d_inner] -> [l_global, b, d_inner//cp] - z = _all_to_all_cp2hp(z, self.cp_group) + if include_gate: + z = _all_to_all_cp2hp(z, self.cp_group) # [l_global//cp, b, d_inner] -> [l_global, b, d_inner//cp] x = _all_to_all_cp2hp(x, self.cp_group) @@ -195,7 +204,7 @@ def pre_conv_ssm( # [l_global//cp, b, nheads] -> [l_global, b, nheads//cp] dt = _all_to_all_cp2hp(dt, self.cp_group) - output = torch.cat([z, x, B, C, dt], dim=-1) + output = torch.cat(([z] if include_gate else []) + [x, B, C, dt], dim=-1) if not self.sequence_is_contiguous: output = _undo_attention_load_balancing(output, self.cp_size, packed_seq_params) diff --git a/megatron/core/ssm/mamba_forest_replay.py b/megatron/core/ssm/mamba_forest_replay.py new file mode 100644 index 00000000000..c9f44c47b2c --- /dev/null +++ b/megatron/core/ssm/mamba_forest_replay.py @@ -0,0 +1,112 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Experimental shared-projection replay with one recurrent scan per forest.""" + +from functools import lru_cache + +import torch +from torch import Tensor + +from megatron.core.models.hybrid.shared_prefix_layout import ( + SharedPrefixForestLayout, + SharedPrefixLayout, +) +from megatron.core.ssm.mamba_mixer import MambaMixer +from megatron.core.ssm.mamba_sequence_packing import scan_mamba_packed_recurrence + + +@lru_cache(maxsize=32) +def forest_replay_indices( + roots: tuple[tuple[int, tuple[int, ...]], ...], device: torch.device +) -> tuple[Tensor, Tensor, Tensor, tuple[int, ...]]: + """Cache expansion, output selection and ordered backward contribution maps.""" + if not roots or any(p < 1 or not cs or min(cs) < 1 for p, cs in roots): + raise ValueError("Forest replay requires positive prefix and completion spans") + physical = sum(p + sum(cs) for p, cs in roots) + expanded = sum(len(cs) * p + sum(cs) for p, cs in roots) + copies = max(len(cs) for _, cs in roots) + dense_indices = [] + output_indices = [] + # The sentinel addresses the appended zero row in the gradient. Each column + # contributes once to a physical prefix token, in reverse sibling order. + contributors = [[expanded] * physical for _ in range(copies)] + lengths = [] + source_offset = dense_offset = 0 + for prefix, completions in roots: + output_indices.extend(range(dense_offset, dense_offset + prefix)) + completion_offset = source_offset + prefix + for branch, length in enumerate(completions): + dense_indices.extend(range(source_offset, source_offset + prefix)) + dense_indices.extend(range(completion_offset, completion_offset + length)) + output_indices.extend(range(dense_offset + prefix, dense_offset + prefix + length)) + column = len(completions) - 1 - branch + contributors[column][source_offset : source_offset + prefix] = range( + dense_offset, dense_offset + prefix + ) + contributors[0][completion_offset : completion_offset + length] = range( + dense_offset + prefix, dense_offset + prefix + length + ) + lengths.append(prefix + length) + completion_offset += length + dense_offset += prefix + length + source_offset = completion_offset + assert source_offset == physical and dense_offset == expanded + assert len(dense_indices) == expanded and len(output_indices) == physical + return ( + torch.tensor(dense_indices, device=device, dtype=torch.long), + torch.tensor(output_indices, device=device, dtype=torch.long), + torch.tensor(contributors, device=device, dtype=torch.long), + tuple(lengths), + ) + + +class _ExpandForest(torch.autograd.Function): + @staticmethod + def forward(ctx, projected, dense_indices, contributors): + """Gather logical dense branch copies from one physical forest.""" + ctx.save_for_backward(contributors) + return projected.index_select(0, dense_indices) + + @staticmethod + def backward(ctx, gradient): + """Accumulate each shared row in a fixed floating-point addition order.""" + (contributors,) = ctx.saved_tensors + # Accumulate shared-prefix contributions in a fixed FP32 order, then + # round once to the model dtype. Avoid BF16 atomic scatter accumulation. + padded = torch.cat([gradient, gradient.new_zeros((1, *gradient.shape[1:]))], dim=0) + accumulation_dtype = torch.float64 if gradient.dtype == torch.float64 else torch.float32 + result = torch.zeros( + (contributors.shape[1], *gradient.shape[1:]), + device=gradient.device, + dtype=accumulation_dtype, + ) + for indices in contributors.unbind(0): + result.add_(padded.index_select(0, indices).to(accumulation_dtype)) + return result.to(gradient.dtype), None, None + + +def scan_mamba_forest_replay( + mixer: MambaMixer, projected: Tensor, layout: SharedPrefixLayout | SharedPrefixForestLayout +) -> Tensor: + """Expand recurrence fields only; run all roots in a single packed scan. + + Prefix projection, gating, attention and MoE remain shared. The recurrent + computation repeats prefixes so that every packed sequence starts from zero, + enabling the ordinary chunk-aligned packed scan without varlen state-fork + backward support. This trades arithmetic for fewer scan calls and removes + the max-sibling rectangle padding used by the state-fork implementation. + """ + if projected.shape[0] < layout.total_len: + raise ValueError("Forest replay input is shorter than its layout") + roots = [] + for offset, root in layout.iter_roots(): + lengths = list(root.completion_lens) + if offset + root.total_len == layout.total_len: + lengths[-1] += projected.shape[0] - layout.total_len + roots.append((root.prefix_len, tuple(lengths))) + dense_indices, output_indices, contributors, lengths = forest_replay_indices( + tuple(roots), projected.device + ) + expanded = _ExpandForest.apply(projected, dense_indices, contributors) + output = scan_mamba_packed_recurrence(mixer, expanded, lengths) + return output.index_select(0, output_indices) diff --git a/megatron/core/ssm/mamba_mixer.py b/megatron/core/ssm/mamba_mixer.py index 0ec5c45625a..37620676a2c 100644 --- a/megatron/core/ssm/mamba_mixer.py +++ b/megatron/core/ssm/mamba_mixer.py @@ -488,7 +488,12 @@ def _mamba_chunk( """Run Mamba through its normalized SSM output, before output projection.""" zxBCdt, _ = self.in_proj(hidden_states) - zxBCdt = self.cp.pre_conv_ssm(zxBCdt, packed_seq_params) + local_gate = None + if not inference_mode and self.use_mem_eff_path and self.config.sequence_relative_kernels: + # Gated norm runs in the original local token/channel order. The + # recurrence never reads z, so avoid its round-trip CP exchange. + local_gate = zxBCdt[..., : self.cp.d_inner_local_tp].contiguous() + zxBCdt = self.cp.pre_conv_ssm(zxBCdt, packed_seq_params, include_gate=local_gate is None) if inference_mode or not self.use_mem_eff_path: # TODO(ksanthanam): Consider deprecating this path for training @@ -499,7 +504,10 @@ def _mamba_chunk( y = self._static_prefill(zxBCdt, conv_state=conv_state, ssm_state=ssm_state) else: assert ssm_state is None - y = self._ssm_training(zxBCdt, packed_seq_params) + if local_gate is None: + y = self._ssm_training(zxBCdt, packed_seq_params) + else: + y = self._ssm_training(zxBCdt, packed_seq_params, local_gate=local_gate) return y @@ -712,7 +720,11 @@ def _static_prefill( return y def _ssm_training( - self, zxBCdt: torch.Tensor, packed_seq_params: Optional[PackedSeqParams] = None + self, + zxBCdt: torch.Tensor, + packed_seq_params: Optional[PackedSeqParams] = None, + *, + local_gate: torch.Tensor | None = None, ) -> torch.Tensor: """ Performs SSM computation for training step. @@ -722,6 +734,17 @@ def _ssm_training( training. """ + if self.config.sequence_relative_kernels: + # Import lazily because the packing helper uses the Mamba kernel symbols above. + from megatron.core.ssm.mamba_sequence_packing import mamba_sequence_relative_scan + + return mamba_sequence_relative_scan( + self, zxBCdt, packed_seq_params, local_gate=local_gate + ) + + if local_gate is not None: + raise ValueError("A local Mamba gate requires sequence-relative kernels") + # transpose: l b pd --> b l pd zxBCdt = rearrange(zxBCdt, "l b d -> b l d").contiguous() diff --git a/megatron/core/ssm/mamba_ragged.py b/megatron/core/ssm/mamba_ragged.py new file mode 100644 index 00000000000..f321ab60701 --- /dev/null +++ b/megatron/core/ssm/mamba_ragged.py @@ -0,0 +1,265 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Chunk-aligned Mamba forests with shared prefix state and ragged branches. + +Only the unaligned prefix tail and the convolution halo are copied to siblings. +Every continuation has its own chunk padding; no longest-sibling rectangle is +materialized. The surrounding projections, gate and distributed layout remain +in the canonical shared-token order. +""" + +import os +from dataclasses import dataclass +from functools import lru_cache + +import numpy as np +import torch +import triton +import triton.language as tl +from causal_conv1d import causal_conv1d_fn +from einops import rearrange + +from megatron.core.models.hybrid.shared_prefix_layout import ( + SharedPrefixForestLayout, + SharedPrefixLayout, +) +from megatron.core.ssm.mamba_mixer import MAMBA_HAS_STATE_DTYPE, MambaMixer + + +@triton.jit +def _gather_forward(X, INDEX, Y, WIDTH: tl.constexpr, BLOCK: tl.constexpr): + row = tl.program_id(0) + col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) + source = tl.load(INDEX + row) + value = tl.load(X + source * WIDTH + col, (source >= 0) & (col < WIDTH), other=0) + tl.store(Y + row * WIDTH + col, value, col < WIDTH) + + +@triton.jit +def _gather_backward( + DY, + CONTRIBUTORS, + DX, + ROWS: tl.constexpr, + WIDTH: tl.constexpr, + COPIES: tl.constexpr, + BLOCK: tl.constexpr, +): + row = tl.program_id(0) + col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) + value = tl.zeros((BLOCK,), tl.float32) + # One owner per input element: deterministic FP32 accumulation, no atomics. + for copy in range(COPIES - 1, -1, -1): + source = tl.load(CONTRIBUTORS + copy * ROWS + row) + value += tl.load(DY + source * WIDTH + col, (source >= 0) & (col < WIDTH), other=0).to( + tl.float32 + ) + tl.store(DX + row * WIDTH + col, value, col < WIDTH) + + +class _RaggedGather(torch.autograd.Function): + @staticmethod + def forward(ctx, value, indices, contributors): + """Gather ragged token rows using the explicit ownership indices.""" + value = value.contiguous() + output = value.new_empty((indices.numel(), 1, value.shape[-1])) + ctx.save_for_backward(contributors) + ctx.input_shape = value.shape + _gather_forward[(indices.numel(), triton.cdiv(value.shape[-1], 256))]( + value, indices, output, value.shape[-1], 256 + ) + return output + + @staticmethod + def backward(ctx, gradient): + """Sum logical copy gradients with one owner for every physical input element.""" + (contributors,) = ctx.saved_tensors + gradient = gradient.contiguous() + output = gradient.new_empty(ctx.input_shape) + _gather_backward[(output.shape[0], triton.cdiv(output.shape[-1], 256))]( + gradient, + contributors, + output, + output.shape[0], + output.shape[-1], + contributors.shape[0], + 256, + ) + return output, None, None + + +@dataclass(frozen=True) +class RaggedMambaLayout: + """Cached token maps and chunk topology; these tensors carry no gradients.""" + + convolution_indices: torch.Tensor + convolution_sequences: torch.Tensor + contributors: torch.Tensor + scan_indices: torch.Tensor + output_indices: torch.Tensor + segment_chunks: torch.Tensor + root_segments: torch.Tensor + input_tokens: int + scan_tokens: int + replayed_tail_tokens: int + padding_tokens: int + + +@lru_cache(maxsize=16) +def ragged_mamba_layout( + roots: tuple[tuple[int, tuple[int, ...]], ...], + chunk_size: int, + convolution_width: int, + device: torch.device, +) -> RaggedMambaLayout: + """Build a forest whose prefix ends at the established scan chunk boundary.""" + if not roots or chunk_size < 1 or chunk_size & (chunk_size - 1): + raise ValueError("Ragged Mamba requires roots and a power-of-two chunk size") + if convolution_width < 1: + raise ValueError("Ragged Mamba requires a positive convolution width") + convolution_indices, convolution_sequences, scan_indices, output_indices = ([], [], [], []) + segment_chunks, root_segments = [0], [0] + input_start = scan_start = replayed = padding = 0 + halo = convolution_width - 1 + + for prefix, lengths in roots: + if prefix < 1 or not lengths or min(lengths) < 0: + raise ValueError("Ragged Mamba requires a nonempty prefix and valid siblings") + head = prefix // chunk_size * chunk_size + tail = prefix - head + segments = [(list(range(input_start, input_start + head)), [-1] * halo)] + branch_start = input_start + prefix + context = [input_start + i if i >= 0 else -1 for i in range(head - halo, head)] + for length in lengths: + tokens = list(range(input_start + head, input_start + prefix)) + tokens.extend(range(branch_start, branch_start + length)) + segments.append((tokens, context)) + branch_start += length + replayed += tail * (len(lengths) - 1) + + for segment, (tokens, context) in enumerate(segments): + real_length = len(tokens) + padded = (real_length + chunk_size - 1) // chunk_size * chunk_size + padding += padded - real_length + if padded: + convolution_start = len(convolution_indices) + convolution_indices.extend(context) + convolution_indices.extend(tokens) + convolution_indices.extend([-1] * (padded - real_length)) + convolution_sequences.extend([len(segment_chunks) - 1] * (halo + padded)) + scan_indices.extend( + range(convolution_start + halo, convolution_start + halo + padded) + ) + if segment == 0: + output_indices.extend(range(scan_start, scan_start + head)) + else: + if segment == 1: + output_indices.extend(range(scan_start, scan_start + tail)) + output_indices.extend(range(scan_start + tail, scan_start + real_length)) + scan_start += padded + segment_chunks.append(scan_start // chunk_size) + root_segments.append(len(segment_chunks) - 1) + input_start = branch_start + + # Preserve the original contributor order without allocating one Python + # list per token or iterating over every token/copy pair in Python. Stable + # sorting keeps destinations in ascending order within each source token; + # the backward kernel still accumulates them in exactly the reverse order. + sources = np.asarray(convolution_indices, dtype=np.int32) + destinations = np.flatnonzero(sources >= 0) + sources = sources[destinations] + order = np.argsort(sources, kind="stable") + counts = np.bincount(sources, minlength=input_start) + starts = np.cumsum(counts) - counts + copy_indices = np.arange(sources.size) - np.repeat(starts, counts) + inverse = np.full((int(counts.max()), input_start), -1, dtype=np.int32) + inverse[copy_indices, sources[order]] = destinations[order] + if len(output_indices) != input_start: + raise RuntimeError("Ragged Mamba output map does not cover the canonical input") + tensor = lambda values: torch.tensor(values, dtype=torch.int32, device=device) + return RaggedMambaLayout( + convolution_indices=tensor(convolution_indices), + convolution_sequences=tensor(convolution_sequences)[None], + contributors=tensor(inverse), + scan_indices=tensor(scan_indices), + output_indices=tensor(output_indices), + segment_chunks=tensor(segment_chunks), + root_segments=tensor(root_segments), + input_tokens=input_start, + scan_tokens=scan_start, + replayed_tail_tokens=replayed, + padding_tokens=padding, + ) + + +def scan_mamba_ragged_forest( + mixer: MambaMixer, + recurrent: torch.Tensor, + layout: SharedPrefixLayout | SharedPrefixForestLayout, +) -> torch.Tensor: + """Scan one canonical forest, retaining autograd through every state fork. + + NRL_SP_MAMBA_SAVE_INTERMEDIATES=1 retains scan intermediates during training. + It defaults to 0; the scan also disables retention for no-grad forwards. + """ + save_intermediates = os.environ.get("NRL_SP_MAMBA_SAVE_INTERMEDIATES", "0") + if save_intermediates not in ("0", "1"): + raise ValueError("NRL_SP_MAMBA_SAVE_INTERMEDIATES must be exactly '0' or '1'") + + from megatron.core.ssm.mamba_ragged_scan import mamba_chunk_scan_forest + + if recurrent.ndim != 3 or recurrent.shape[1] != 1: + raise ValueError("Ragged Mamba requires [sequence, 1, recurrent channels]") + if recurrent.shape[0] < layout.total_len: + raise ValueError("Ragged Mamba input is shorter than its forest") + roots = [] + for offset, root in layout.iter_roots(): + lengths = list(root.completion_lens) + if offset + root.total_len == layout.total_len: + lengths[-1] += recurrent.shape[0] - layout.total_len + roots.append((root.prefix_len, tuple(lengths))) + metadata = ragged_mamba_layout(tuple(roots), mixer.chunk_size, mixer.d_conv, recurrent.device) + cp = mixer.cp + dim, groups = cp.d_inner_local_tpcp, cp.ngroups_local_tpcp + fields = _RaggedGather.apply( + recurrent, metadata.convolution_indices, metadata.contributors + ).transpose(0, 1) + xbc, dt = fields.split([dim + 2 * groups * mixer.d_state, cp.nheads_local_tpcp], -1) + xbc = ( + causal_conv1d_fn( + xbc.contiguous().transpose(1, 2), + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + seq_idx=metadata.convolution_sequences, + activation=mixer.activation, + ) + .transpose(1, 2) + .index_select(1, metadata.scan_indices) + .contiguous() + ) + dt = dt.index_select(1, metadata.scan_indices).contiguous() + x, b, c = xbc.split([dim, groups * mixer.d_state, groups * mixer.d_state], -1) + # The scan kernels accept these strided views: each field's last dimension + # is contiguous. Retain the shared xbc storage rather than copying fields. + y = mamba_chunk_scan_forest( + rearrange(x, "b l (h p) -> b l h p", p=mixer.headdim), + dt, + -torch.exp(cp.get_A_log().float()), + rearrange(b, "b l (g n) -> b l g n", n=mixer.d_state), + rearrange(c, "b l (g n) -> b l g n", n=mixer.d_state), + mixer.chunk_size, + metadata.segment_chunks, + metadata.root_segments, + D=( + rearrange(cp.get_D().float(), "(h p) -> h p", p=mixer.headdim) + if mixer.D_has_hdim + else cp.get_D() + ), + dt_bias=cp.get_dt_bias().float(), + dt_softplus=True, + state_dtype=mixer.mamba_training_ssm_states_dtype if MAMBA_HAS_STATE_DTYPE else None, + save_intermediates=save_intermediates == "1" and mixer.training, + ) + return ( + rearrange(y, "b l h p -> l b (h p)").contiguous().index_select(0, metadata.output_indices) + ) diff --git a/megatron/core/ssm/mamba_ragged_scan.py b/megatron/core/ssm/mamba_ragged_scan.py new file mode 100644 index 00000000000..0e6e92c8e7a --- /dev/null +++ b/megatron/core/ssm/mamba_ragged_scan.py @@ -0,0 +1,498 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2024, Tri Dao, Albert Gu. +# +# The combined-scan forward/backward below adapts mamba_ssm's ssd_combined.py, +# and state passing adapts ssd_state_passing.py. The upstream portions are +# licensed under Apache-2.0; see the Mamba license in the repository LICENSE. + +"""Differentiable SSD scan over chunk-aligned shared-prefix forests. + +Only inter-chunk state passing changes topology. All intra-chunk operations +reuse the installed mamba_ssm kernels. The first segment of each root is its +prefix; its remaining segments are siblings, each initialized from the prefix +terminal state. No prefix recurrence is replayed and no sibling is padded to +another sibling's length. +""" + +import torch +import triton +import triton.language as tl +from einops import rearrange +from mamba_ssm.ops.triton.ssd_bmm import _bmm_chunk_bwd, _bmm_chunk_fwd +from mamba_ssm.ops.triton.ssd_chunk_scan import ( + _chunk_scan_bwd_dC, + _chunk_scan_bwd_dcb, + _chunk_scan_bwd_ddAcs_stable, + _chunk_scan_bwd_dstates, + _chunk_scan_fwd, +) +from mamba_ssm.ops.triton.ssd_chunk_state import ( + _chunk_cumsum_bwd, + _chunk_cumsum_fwd, + _chunk_state_bwd_db, + _chunk_state_fwd, +) +from mamba_ssm.ops.triton.ssd_combined import _chunk_scan_chunk_state_bwd_dx +from torch import Tensor +from torch.autograd.function import once_differentiable + +_STATE_BLOCK = 256 + + +@triton.jit +def _forest_state_fwd_kernel( + chunk_states, + dA, + segment_chunks, + root_segments, + out, + dim: tl.constexpr, + stride_chunk: tl.constexpr, + stride_head: tl.constexpr, + stride_dim: tl.constexpr, + stride_a_head: tl.constexpr, + stride_a_chunk: tl.constexpr, + stride_out_chunk: tl.constexpr, + stride_out_head: tl.constexpr, + BLOCK: tl.constexpr, +): + tile = tl.program_id(0) + root = tl.program_id(1) + head = tl.program_id(2) + offsets = tile * BLOCK + tl.arange(0, BLOCK) + valid = offsets < dim + state_base = chunk_states + head * stride_head + offsets * stride_dim + out_base = out + head * stride_out_head + offsets + a_base = dA + head * stride_a_head + first_segment = tl.load(root_segments + root) + stop_segment = tl.load(root_segments + root + 1) + prefix_start = tl.load(segment_chunks + first_segment) + prefix_stop = tl.load(segment_chunks + first_segment + 1) + + # Keep the FP32 accumulator live across chunks, including at the fork. + # The stored entries have the same rounding as upstream state passing. + prefix_state = tl.full((BLOCK,), 0.0, tl.float32) + for chunk in range(prefix_start, prefix_stop): + tl.store(out_base + chunk * stride_out_chunk, prefix_state, valid) + delta = tl.load(state_base + chunk * stride_chunk, valid, 0.0).to(tl.float32) + scale = tl.exp(tl.load(a_base + chunk * stride_a_chunk).to(tl.float32)) + prefix_state = scale * prefix_state + delta + + for segment in range(first_segment + 1, stop_segment): + start = tl.load(segment_chunks + segment) + stop = tl.load(segment_chunks + segment + 1) + state = prefix_state + for chunk in range(start, stop): + tl.store(out_base + chunk * stride_out_chunk, state, valid) + delta = tl.load(state_base + chunk * stride_chunk, valid, 0.0).to(tl.float32) + scale = tl.exp(tl.load(a_base + chunk * stride_a_chunk).to(tl.float32)) + state = scale * state + delta + + +@triton.jit +def _forest_state_bwd_kernel( + entries, + direct_grads, + dA, + segment_chunks, + root_segments, + chunk_grads, + dA_tiles, + converted_entries, + dim: tl.constexpr, + stride_entry_chunk: tl.constexpr, + stride_entry_head: tl.constexpr, + stride_entry_dim: tl.constexpr, + stride_direct_chunk: tl.constexpr, + stride_direct_head: tl.constexpr, + stride_direct_dim: tl.constexpr, + stride_a_head: tl.constexpr, + stride_a_chunk: tl.constexpr, + stride_grad_chunk: tl.constexpr, + stride_grad_head: tl.constexpr, + stride_da_head: tl.constexpr, + stride_da_chunk: tl.constexpr, + stride_converted_chunk: tl.constexpr, + stride_converted_head: tl.constexpr, + CONVERT_ENTRIES: tl.constexpr, + BLOCK: tl.constexpr, +): + tile = tl.program_id(0) + root = tl.program_id(1) + head = tl.program_id(2) + offsets = tile * BLOCK + tl.arange(0, BLOCK) + valid = offsets < dim + entry_base = entries + head * stride_entry_head + offsets * stride_entry_dim + direct_base = direct_grads + head * stride_direct_head + offsets * stride_direct_dim + grad_base = chunk_grads + head * stride_grad_head + offsets + a_base = dA + head * stride_a_head + da_base = dA_tiles + head * stride_da_head + tile + if CONVERT_ENTRIES: + converted_base = converted_entries + head * stride_converted_head + offsets + first_segment = tl.load(root_segments + root) + stop_segment = tl.load(root_segments + root + 1) + prefix_start = tl.load(segment_chunks + first_segment) + prefix_stop = tl.load(segment_chunks + first_segment + 1) + + # Sum siblings in a fixed reverse order in FP32. No atomic state updates, + # no materialized per-branch initial-state gradients, and no host loops. + prefix_adjoint = tl.full((BLOCK,), 0.0, tl.float32) + for segment in range(stop_segment - 1, first_segment, -1): + start = tl.load(segment_chunks + segment) + stop = tl.load(segment_chunks + segment + 1) + adjoint = tl.full((BLOCK,), 0.0, tl.float32) + for chunk in range(stop - 1, start - 1, -1): + tl.store(grad_base + chunk * stride_grad_chunk, adjoint, valid) + entry = tl.load(entry_base + chunk * stride_entry_chunk, valid, 0.0).to(tl.float32) + scale = tl.exp(tl.load(a_base + chunk * stride_a_chunk).to(tl.float32)) + # This term includes the first branch chunk: its initial state is + # the prefix terminal, so dropping it would lose dA at the fork. + tl.store(da_base + chunk * stride_da_chunk, tl.sum(entry * adjoint) * scale) + direct = tl.load(direct_base + chunk * stride_direct_chunk, valid, 0.0).to(tl.float32) + adjoint = scale * adjoint + direct + if CONVERT_ENTRIES: + tl.store(converted_base + chunk * stride_converted_chunk, entry, valid) + prefix_adjoint += adjoint + + for chunk in range(prefix_stop - 1, prefix_start - 1, -1): + tl.store(grad_base + chunk * stride_grad_chunk, prefix_adjoint, valid) + entry = tl.load(entry_base + chunk * stride_entry_chunk, valid, 0.0).to(tl.float32) + scale = tl.exp(tl.load(a_base + chunk * stride_a_chunk).to(tl.float32)) + tl.store(da_base + chunk * stride_da_chunk, tl.sum(entry * prefix_adjoint) * scale) + direct = tl.load(direct_base + chunk * stride_direct_chunk, valid, 0.0).to(tl.float32) + prefix_adjoint = scale * prefix_adjoint + direct + if CONVERT_ENTRIES: + tl.store(converted_base + chunk * stride_converted_chunk, entry, valid) + + +def _forest_state_fwd(states, dA, segment_chunks, root_segments, out_dtype): + _, nchunks, nheads, dim = states.shape + out = torch.empty((1, nchunks, nheads, dim), device=states.device, dtype=out_dtype) + grid = (triton.cdiv(dim, _STATE_BLOCK), root_segments.numel() - 1, nheads) + with torch.cuda.device(states.device.index): + _forest_state_fwd_kernel[grid]( + states, + dA, + segment_chunks, + root_segments, + out, + dim, + states.stride(1), + states.stride(2), + states.stride(3), + dA.stride(1), + dA.stride(2), + out.stride(1), + out.stride(2), + BLOCK=_STATE_BLOCK, + ) + return out + + +def _forest_state_bwd(entries, dA, direct_grads, segment_chunks, root_segments, dtype): + _, nchunks, nheads, dim = entries.shape + chunk_grads = torch.empty((1, nchunks, nheads, dim), device=entries.device, dtype=dtype) + converted = entries if entries.dtype == dtype else torch.empty_like(chunk_grads) + tiles = triton.cdiv(dim, _STATE_BLOCK) + dA_tiles = torch.empty((1, nheads, nchunks, tiles), device=dA.device, dtype=torch.float32) + grid = (tiles, root_segments.numel() - 1, nheads) + with torch.cuda.device(entries.device.index): + _forest_state_bwd_kernel[grid]( + entries, + direct_grads, + dA, + segment_chunks, + root_segments, + chunk_grads, + dA_tiles, + converted, + dim, + entries.stride(1), + entries.stride(2), + entries.stride(3), + direct_grads.stride(1), + direct_grads.stride(2), + direct_grads.stride(3), + dA.stride(1), + dA.stride(2), + chunk_grads.stride(1), + chunk_grads.stride(2), + dA_tiles.stride(1), + dA_tiles.stride(2), + converted.stride(1), + converted.stride(2), + CONVERT_ENTRIES=converted is not entries, + BLOCK=_STATE_BLOCK, + ) + return chunk_grads, dA_tiles.sum(dim=-1).to(dA.dtype), converted + + +def _forest_scan_fwd( + x, + dt, + A, + B, + C, + chunk_size, + segment_chunks, + root_segments, + D, + dt_bias, + dt_softplus, + save_intermediates=False, +): + dA_cumsum, dt_out = _chunk_cumsum_fwd( + dt, A, chunk_size, dt_bias=dt_bias, dt_softplus=dt_softplus + ) + chunk_states = _chunk_state_fwd(B, x, dt_out, dA_cumsum, states_in_fp32=True) + states = _forest_state_fwd( + chunk_states.flatten(-2), dA_cumsum[..., -1], segment_chunks, root_segments, C.dtype + ).unflatten(-1, x.shape[-1:] + B.shape[-1:]) + if not save_intermediates: + del chunk_states + CB = _bmm_chunk_fwd(C, B, chunk_size, output_dtype=torch.float32) + out, _ = _chunk_scan_fwd(CB, x, dt_out, dA_cumsum, C, states, D=D, z=None) + intermediates = (dA_cumsum, dt_out, CB, chunk_states) if save_intermediates else () + return out, intermediates + + +def _forest_scan_bwd( + dout, + x, + dt_in, + A, + B, + C, + chunk_size, + segment_chunks, + root_segments, + D, + dt_bias, + dt_softplus, + intermediates=(), +): + # The default recomputes all intermediates, as the pinned combined backward + # does. Opt-in saved tensors are the exact FP32 forward results. In both + # paths chunk-entry states are reconstructed in FP32 below; the rounded + # forward entry states are never substituted for this backward recurrence. + if dout.stride(-1) != 1: + dout = dout.contiguous() + dt_in = dt_in.clone() # Preserve the upstream Triton device-context workaround. + if intermediates: + dA_cumsum, dt, CB, chunk_states = intermediates + else: + dA_cumsum, dt = _chunk_cumsum_fwd( + dt_in, A, chunk_size, dt_bias=dt_bias, dt_softplus=dt_softplus + ) + CB = _bmm_chunk_fwd(C, B, chunk_size, output_dtype=torch.float32) + chunk_states = _chunk_state_fwd(B, x, dt, dA_cumsum, states_in_fp32=True) + states = _forest_state_fwd( + chunk_states.flatten(-2), dA_cumsum[..., -1], segment_chunks, root_segments, torch.float32 + ).unflatten(-1, x.shape[-1:] + B.shape[-1:]) + del chunk_states + direct = _chunk_scan_bwd_dstates(C, dA_cumsum, dout, dtype=states.dtype) + dstates, ddA_chunk_cumsum, states = _forest_state_bwd( + states.flatten(-2), + dA_cumsum[..., -1], + direct.flatten(-2), + segment_chunks, + root_segments, + x.dtype, + ) + del direct + states = states.unflatten(-1, x.shape[-1:] + B.shape[-1:]) + dstates = dstates.unflatten(-1, x.shape[-1:] + B.shape[-1:]) + ngroups = B.shape[2] + dx, ddt, dD = _chunk_scan_chunk_state_bwd_dx(x, dt, dA_cumsum, B, CB, dout, dstates, D=D) + dB, ddA_next = _chunk_state_bwd_db(x, dt, dA_cumsum, dstates, B=B, ngroups=ngroups) + dC, ddA_cumsum_prev = _chunk_scan_bwd_dC(states, dA_cumsum, dout, C=C, ngroups=ngroups) + dCB = _chunk_scan_bwd_dcb(x, dt, dA_cumsum, dout, ngroups=ngroups).to(CB.dtype) + dB_out, dC_out = torch.empty_like(B), torch.empty_like(C) + _bmm_chunk_bwd(C, dCB, residual=dB, out=dB_out) + _bmm_chunk_bwd(B, rearrange(dCB, "... l s -> ... s l"), residual=dC, out=dC_out) + ddA_cumsum_prev[..., -1] += ddA_chunk_cumsum + ddA_prev = ddA_cumsum_prev.flip([-1]).cumsum(dim=-1).flip([-1]) + ddA = _chunk_scan_bwd_ddAcs_stable(x, dt, dA_cumsum, dout, CB) + ddA += ddA_next + ddA_prev + ddt_out, dA, ddt_bias = _chunk_cumsum_bwd( + ddA, ddt, dt_in, A, dt_bias=dt_bias, dt_softplus=dt_softplus + ) + return dx, ddt_out, dA, dB_out, dC_out, dD, ddt_bias + + +class _MambaChunkScanForest(torch.autograd.Function): + @staticmethod + def forward( + ctx, + x, + dt, + A, + B, + C, + chunk_size, + segment_chunks, + root_segments, + D, + dt_bias, + dt_softplus, + save_intermediates, + ): + """Run the chunk-aligned forest scan and retain inputs needed by backward.""" + if B.stride(-1) != 1: + B = B.contiguous() + if C.stride(-1) != 1: + C = C.contiguous() + if x.stride(-1) != 1 and x.stride(1) != 1: + x = x.contiguous() + if D is not None and D.stride(-1) != 1: + D = D.contiguous() + out, intermediates = _forest_scan_fwd( + x, + dt, + A, + B, + C, + chunk_size, + segment_chunks, + root_segments, + D, + dt_bias, + dt_softplus, + save_intermediates=save_intermediates, + ) + ctx.save_for_backward( + x, dt, A, B, C, segment_chunks, root_segments, D, dt_bias, *intermediates + ) + ctx.chunk_size = chunk_size + ctx.dt_softplus = dt_softplus + return out + + @staticmethod + @once_differentiable + def backward(ctx, dout): + """Propagate output gradients through forked scan states and chunk parameters.""" + x, dt, A, B, C, segment_chunks, root_segments, D, dt_bias, *intermediates = ( + ctx.saved_tensors + ) + dx, ddt, dA, dB, dC, dD, ddt_bias = _forest_scan_bwd( + dout, + x, + dt, + A, + B, + C, + ctx.chunk_size, + segment_chunks, + root_segments, + D, + dt_bias, + ctx.dt_softplus, + intermediates=intermediates, + ) + return dx, ddt, dA, dB, dC, None, None, None, dD, ddt_bias, None, None + + +def mamba_chunk_scan_forest( + x: Tensor, + dt: Tensor, + A: Tensor, + B: Tensor, + C: Tensor, + chunk_size: int, + segment_chunks: Tensor, + root_segments: Tensor, + *, + D: Tensor | None = None, + dt_bias: Tensor | None = None, + dt_softplus: bool = True, + state_dtype: torch.dtype | None = None, + save_intermediates: bool = False, +) -> Tensor: + """Scan independent shared-prefix roots with differentiable state forks. + + Args: + x: Input with shape ``[1, T, H, P]``; T must be chunk aligned. + dt: Time steps with shape ``[1, T, H]``. + A: Decay parameters with shape ``[H]``. + B: Input state projections with shape ``[1, T, G, N]``. + C: Output state projections with the same shape as B. + chunk_size: Positive power-of-two chunk size supported by mamba_ssm. + segment_chunks: Contiguous CUDA int32 cumulative chunk boundaries, + starting at zero and ending at ``T // chunk_size``. Monotone + nondecreasing boundaries allow empty segments. + root_segments: Contiguous CUDA int32 cumulative segment boundaries, + starting at zero and ending at ``segment_chunks.numel() - 1``. + Boundaries must increase strictly. The first segment of each root + is its prefix; all later segments are independent siblings. + D: Optional skip weights with shape ``[H]`` or ``[H, P]``. + dt_bias: Optional time-step bias with shape ``[H]``. + dt_softplus: Apply softplus to biased time steps. + state_dtype: Only the pinned upstream default (None or C.dtype) is + supported. Inter-chunk accumulators always remain FP32. + save_intermediates: Opt in to retaining the FP32 time-step/cumulative + decay, C-B products and chunk states for backward. This removes + three recomputation kernels at the cost of four saved tensors. + Defaults to recomputation and is disabled when grad mode is off + or no floating input needs a gradient. State passing is unchanged. + + Returns: + Output with shape and dtype matching x, including supplied tail padding. + + Notes: + The caller must validate the metadata values when constructing the + cached layout. This function validates tensor metadata without device + synchronization or per-layer GPU validation launches. All segments + must start on chunk boundaries; convolution context is the caller's + responsibility. No gate, external states, or second derivatives are + provided. Padding must be excluded from downstream loss/output maps. + """ + if x.ndim != 4 or x.shape[0] != 1 or x.shape[1] == 0: + raise ValueError("Forest scan requires nonempty x with shape [1, T, H, P]") + if chunk_size < 1 or chunk_size & (chunk_size - 1) or x.shape[1] % chunk_size: + raise ValueError("Forest scan requires power-of-two chunk size and chunk-aligned T") + _, length, heads, headdim = x.shape + if B.ndim != 4 or B.shape[:2] != (1, length) or C.shape != B.shape: + raise ValueError("Forest B and C must have matching shape [1, T, G, N]") + if B.shape[2] < 1 or B.shape[3] < 1 or heads < 1 or headdim < 1 or heads % B.shape[2]: + raise ValueError("Forest state dimensions must be positive and groups must divide heads") + if dt.shape != (1, length, heads) or A.shape != (heads,): + raise ValueError("Forest dt and A dimensions must match x heads and tokens") + if D is not None and D.shape not in ((heads,), (heads, headdim)): + raise ValueError("Forest D must have shape [H] or [H, P]") + if dt_bias is not None and dt_bias.shape != (heads,): + raise ValueError("Forest dt_bias must have shape [H]") + tensors = (x, dt, A, B, C, segment_chunks, root_segments, D, dt_bias) + if not x.is_cuda or any(t.device != x.device for t in tensors if t is not None): + raise ValueError("Forest scan tensors must reside on one CUDA device") + for boundaries in (segment_chunks, root_segments): + if ( + boundaries.ndim != 1 + or boundaries.numel() < 2 + or boundaries.dtype != torch.int32 + or not boundaries.is_contiguous() + ): + raise ValueError("Forest boundaries must be contiguous int32 vectors of length >= 2") + if state_dtype not in (None, C.dtype): + raise NotImplementedError("Forest scan currently requires the pinned upstream state dtype") + if not isinstance(save_intermediates, bool): + raise TypeError("save_intermediates must be a bool") + if save_intermediates: + # Uniform checkpoint's first forward runs under no_grad; retain only + # during its grad-enabled backward replay, not across checkpoint spans. + save_intermediates = torch.is_grad_enabled() and any( + tensor.requires_grad for tensor in (x, dt, A, B, C, D, dt_bias) if tensor is not None + ) + return _MambaChunkScanForest.apply( + x, + dt, + A, + B, + C, + chunk_size, + segment_chunks, + root_segments, + D, + dt_bias, + dt_softplus, + save_intermediates, + ) diff --git a/megatron/core/ssm/mamba_sequence_packing.py b/megatron/core/ssm/mamba_sequence_packing.py new file mode 100644 index 00000000000..4f591ac0b81 --- /dev/null +++ b/megatron/core/ssm/mamba_sequence_packing.py @@ -0,0 +1,136 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Experimental packed Mamba with a sequence-relative scan chunk grid.""" + +from functools import lru_cache +from itertools import accumulate + +import torch +from causal_conv1d import causal_conv1d_fn +from einops import rearrange + +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.ssm.mamba_mixer import ( + MAMBA_HAS_STATE_DTYPE, + MambaMixer, + mamba_chunk_scan_combined, +) + + +@lru_cache(maxsize=32) +def aligned_indices( + lengths: tuple[int, ...], chunk_size: int, device: torch.device +) -> tuple[torch.Tensor, torch.Tensor, int]: + """Cache bounded, non-differentiable layout metadata across model layers.""" + padded = tuple((n + chunk_size - 1) // chunk_size * chunk_size for n in lengths) + starts = tuple(accumulate(padded, initial=0)) + real = torch.cat( + [ + torch.arange(start, start + n, device=device) + for start, n in zip(starts[:-1], lengths, strict=True) + ] + ) + sequence_ids = torch.repeat_interleave( + torch.arange(len(lengths), device=device, dtype=torch.int32), + torch.tensor(padded, device=device), + output_size=starts[-1], + )[None] + return real, sequence_ids, starts[-1] + + +def mamba_sequence_relative_scan( + mixer: MambaMixer, + projected: torch.Tensor, + packed_seq_params: PackedSeqParams | None = None, + *, + local_gate: torch.Tensor | None = None, +) -> torch.Tensor: + """Pad only the internal recurrence layout, then restore original CP positions. + + Every sequence begins on a scan chunk boundary. The convolution still uses + sequence IDs to reset its causal context. Tail padding has no path back to + real outputs; attention and the surrounding model retain the original layout. + This is a new numerical reference, not an emulation of stock packed Mamba. + """ + if not mixer.rmsnorm or mixer.norm_before_gate: + raise ValueError("sequence-relative Mamba requires RMSNorm after gating") + if projected.ndim != 3 or projected.shape[1] != 1: + raise ValueError("sequence-relative Mamba requires packed [sequence, 1, projection] input") + if packed_seq_params is None: + lengths = (projected.shape[0],) + else: + assert packed_seq_params.qkv_format == "thd" + cu = packed_seq_params.cu_seqlens_q_padded + if cu is None: + cu = packed_seq_params.cu_seqlens_q + lengths = tuple((cu[1:] - cu[:-1]).tolist()) + assert sum(lengths) == projected.shape[0] and min(lengths) > 0 + cp = mixer.cp + dim = cp.d_inner_local_tpcp + if local_gate is None: + gate = projected[..., :dim] + recurrent = projected[..., dim:] + else: + gate = local_gate + recurrent = projected + y = scan_mamba_packed_recurrence(mixer, recurrent, lengths) + y = cp.post_conv_ssm(y, packed_seq_params) + if local_gate is None: + gate = cp.post_conv_ssm(gate.contiguous(), packed_seq_params) + if gate.shape != y.shape: + raise ValueError("Mamba gate must match the restored local output layout") + return mixer.norm(y, gate) + + +def scan_mamba_packed_recurrence( + mixer: MambaMixer, recurrent: torch.Tensor, lengths: tuple[int, ...] +) -> torch.Tensor: + """Scan canonical x/B/C/dt sequences with one aligned convolution and scan. + + Returns canonical y without a gate, normalization or CP transform. Each + complete sequence starts at recurrence state zero and a chunk boundary. + """ + if recurrent.ndim != 3 or recurrent.shape[1] != 1: + raise ValueError("Packed recurrence requires [sequence, 1, channels]") + if not lengths or min(lengths) < 1 or sum(lengths) != recurrent.shape[0]: + raise ValueError("Packed recurrence lengths must cover its input") + cp = mixer.cp + dim, groups = cp.d_inner_local_tpcp, cp.ngroups_local_tpcp + real, sequence_ids, total = aligned_indices(lengths, mixer.chunk_size, recurrent.device) + # Fresh storage has no aliases; an in-place copy avoids a second full-sized + # padding allocation while retaining the differentiable source gather. + padded = recurrent.new_zeros((total, 1, recurrent.shape[-1])).index_copy_(0, real, recurrent) + fields = rearrange(padded, "l b d -> b l d").contiguous() + xbc, dt = fields.split([dim + 2 * groups * mixer.d_state, cp.nheads_local_tpcp], -1) + xbc = causal_conv1d_fn( + rearrange(xbc.contiguous(), "b l d -> b d l"), + rearrange(cp.get_conv1d_weight(), "d 1 w -> d w"), + cp.get_conv1d_bias(), + seq_idx=sequence_ids, + activation=mixer.activation, + ) + x, b, c = ( + rearrange(xbc, "b d l -> b l d") + .contiguous() + .split([dim, groups * mixer.d_state, groups * mixer.d_state], -1) + ) + y = mamba_chunk_scan_combined( + rearrange(x, "b l (h p) -> b l h p", p=mixer.headdim).contiguous(), + dt.contiguous(), + -torch.exp(cp.get_A_log().float()), + rearrange(b, "b l (g n) -> b l g n", n=mixer.d_state).contiguous(), + rearrange(c, "b l (g n) -> b l g n", n=mixer.d_state).contiguous(), + mixer.chunk_size, + D=( + rearrange(cp.get_D().float(), "(h p) -> h p", p=mixer.headdim) + if mixer.D_has_hdim + else cp.get_D() + ), + z=None, + seq_idx=sequence_ids, + dt_bias=cp.get_dt_bias().float(), + dt_softplus=True, + **({"state_dtype": mixer.mamba_training_ssm_states_dtype} if MAMBA_HAS_STATE_DTYPE else {}), + ) + y = rearrange(y, "b l h p -> l b (h p)").contiguous().index_select(0, real) + return y diff --git a/megatron/core/tensor_parallel/ordered_reduce_scatter.py b/megatron/core/tensor_parallel/ordered_reduce_scatter.py new file mode 100644 index 00000000000..94c4239640d --- /dev/null +++ b/megatron/core/tensor_parallel/ordered_reduce_scatter.py @@ -0,0 +1,93 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Sequence-parallel reduction with a fixed FP32 source-rank addition order.""" + +import torch + +try: + import triton + import triton.language as tl +except ImportError: + triton = None + tl = None + + +if triton is not None: + + @triton.jit + def _sum_sources(inp, out, count: tl.constexpr, ranks: tl.constexpr, block: tl.constexpr): + offsets = tl.program_id(0) * block + tl.arange(0, block) + mask = offsets < count + result = tl.load(inp + offsets, mask=mask, other=0).to(tl.float32) + for rank in tl.static_range(1, ranks): + result = result + tl.load(inp + rank * count + offsets, mask=mask, other=0).to( + tl.float32 + ) + tl.store(out + offsets, result, mask=mask) + + +class _OrderedReduceScatter(torch.autograd.Function): + @staticmethod + def forward(ctx, input_: torch.Tensor, group: torch.distributed.ProcessGroup): + """Reduce sequence shards in a fixed source-rank order using FP32 additions.""" + ctx.group = group + size = group.size() + if size == 1: + return input_ + received = torch.empty_like(input_, memory_format=torch.contiguous_format) + torch.distributed.all_to_all_single(received, input_.contiguous(), group=group) + shape = (input_.shape[0] // size, *input_.shape[1:]) + if triton is not None and input_.is_cuda: + output = torch.empty(shape, dtype=input_.dtype, device=input_.device) + count = output.numel() + if count: + # One kernel keeps intermediate FP32 sums in registers. Disable + # fusion and preserve the validated source-rank addition order. + _sum_sources[(triton.cdiv(count, 1024),)]( + received, output, count, size, 1024, enable_fp_fusion=False, num_warps=4 + ) + return output + parts = received.reshape(size, -1) + # NCCL's reduction tree may change with message size. FP32 NCCL SUM + # alone still differs at rare BF16 rounding ties. Communicate in the + # original dtype, then add sources in the same order for every token. + total = parts[0].to(torch.float32, copy=True) + for rank in range(1, size): + total.add_(parts[rank].float()) + return total.reshape(shape).to(input_.dtype) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + """Gather each rank's output gradient into the original unsharded shape.""" + size = ctx.group.size() + if size == 1: + return grad_output, None + shape = (grad_output.shape[0] * size, *grad_output.shape[1:]) + grad_input = torch.empty(shape, dtype=grad_output.dtype, device=grad_output.device) + torch.distributed.all_gather_into_tensor( + grad_input, grad_output.contiguous(), group=ctx.group + ) + return grad_input, None + + +def ordered_reduce_scatter_to_sequence_parallel_region( + input_: torch.Tensor, *, group: torch.distributed.ProcessGroup +) -> torch.Tensor: + """Sum TP partials in source-rank order and shard the leading dimension. + + Communication retains the input dtype; local additions use FP32. Backward + gathers sequence shards without a reduction. This synchronous path requires + extra scratch storage and does not overlap communication with the GEMM. + + Args: + input_: Equal-sized FP16, BF16, or FP32 partial output on every TP rank. + group: Explicit tensor-parallel process group. + + Returns: + This rank's contiguous sequence shard, in the input dtype. + """ + if input_.dtype not in (torch.float16, torch.bfloat16, torch.float32): + raise ValueError("ordered TP reduction requires FP16, BF16, or FP32 inputs") + if input_.ndim == 0 or input_.shape[0] % group.size(): + raise ValueError("the leading dimension must be divisible by the TP group size") + return _OrderedReduceScatter.apply(input_, group) diff --git a/megatron/core/tensor_parallel/random.py b/megatron/core/tensor_parallel/random.py index e6931715c2a..ca2a69afb23 100644 --- a/megatron/core/tensor_parallel/random.py +++ b/megatron/core/tensor_parallel/random.py @@ -7,6 +7,7 @@ import contextlib import logging from collections.abc import Callable +from contextvars import copy_context from typing import Any, Optional, TypeVar, Union import torch @@ -106,6 +107,14 @@ def _get_share_storage(): _EXPERT_GTP_REMAT_RNG_TRACKER_NAME = 'egtp-remat-rng' +def _run_recompute_with_observation_suspended(function, *args): + """Suppress observations inside the restored forward context, not outside it.""" + # Tensor-observation suspension is itself a ContextVar. Entering the copied + # forward context after suspending would restore its old unsuspended value. + with suspend_tensor_observations(): + return function(*args) + + def _get_cuda_rng_state( device: Union[int, str, torch.device] = "cuda", clone: bool = False, graph_safe: bool = False ) -> torch.Tensor: @@ -652,6 +661,10 @@ def forward( _set_checkpointing() ctx.run_function = run_function + # CUDA autograd can recompute on an engine thread whose Python context + # differs from the caller. Preserve forward execution scopes (including + # router GEMM row blocks) instead of inheriting that thread's defaults. + ctx.forward_context = copy_context() ctx.distribute_saved_activations = distribute_saved_activations # Copy the rng states. @@ -699,8 +712,10 @@ def backward(ctx, *args): # Compute the forward pass. detached_inputs = detach_variable(inputs) - with torch.enable_grad(), suspend_tensor_observations(): - outputs = ctx.run_function(*detached_inputs) + with torch.enable_grad(): + outputs = ctx.forward_context.copy().run( + _run_recompute_with_observation_suspended, ctx.run_function, *detached_inputs + ) if isinstance(outputs, torch.Tensor): outputs = (outputs,) @@ -928,6 +943,7 @@ def __init__(self, fp8=False, ckpt_manager=None, retain_input_tensors=False): self.ckpt_manager = ckpt_manager self.retain_input_tensors = retain_input_tensors self.run_function = None + self.forward_context = None self.fwd_cpu_rng_state = None self.fwd_cuda_rng_state = None self.fwd_cuda_rng_state_tracker = None @@ -950,6 +966,7 @@ def checkpoint(self, run_function: Callable[[Unpack[_Ts]], _R], *args: Unpack[_T return run_function(*args) self.run_function = run_function + self.forward_context = copy_context() self.rng_states = _get_all_rng_states() @@ -999,10 +1016,13 @@ def _recompute(self, _): # Reconstruct full args list from saved ctx inputs = _load_args_from_ctx(self.ctx) - with torch.enable_grad(), fp8_ctx, recompute_ctx, suspend_tensor_observations(): - outputs = self.run_function(*inputs) + with torch.enable_grad(), fp8_ctx, recompute_ctx: + outputs = self.forward_context.copy().run( + _run_recompute_with_observation_suspended, self.run_function, *inputs + ) self.run_function = None + self.forward_context = None self.rng_states = None if isinstance(outputs, torch.Tensor): diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index 39ee90d2939..b3595184112 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -1617,7 +1617,50 @@ def forward_pre_attn_and_core_attn( core_attn_manager = off_interface( self.offload_core_attention and self.training, query, "core_attn" ) - if self.checkpoint_core_attention and self.training: + shared_prefix_forest = getattr(self, "_shared_prefix_forest", None) + if shared_prefix_forest is not None: + if inference_context is not None or packed_seq_params is not None: + raise NotImplementedError( + "shared-prefix fused attention only supports static, unpacked training" + ) + if attention_mask is not None or attention_bias is not None: + raise ValueError( + 'shared-prefix fused attention owns its tree mask and does not accept a ' + 'mask/bias' + ) + if query.dtype not in (torch.float16, torch.bfloat16): + raise TypeError("shared-prefix fused attention requires fp16 or bf16 Q/K/V") + + from megatron.core.models.hybrid.shared_prefix_fused import ( + flash_composed_forest_attention, + flash_composed_forest_attention_cp, + ) + + softmax_scale = self.config.softmax_scale or query.shape[-1] ** -0.5 + with core_attn_manager as query: + if self.pg_collection.cp.size() > 1: + core_attn_out = flash_composed_forest_attention_cp( + query, + key, + value, + shared_prefix_forest, + cp_group=self.pg_collection.cp, + scale=softmax_scale, + full_context=self.config.sequence_relative_kernels, + ) + else: + core_attn_out = flash_composed_forest_attention( + query, + key, + value, + shared_prefix_forest, + scale=softmax_scale, + full_context=self.config.sequence_relative_kernels, + ) + core_attn_out = core_attn_manager.group_offload( + core_attn_out, forced_released_tensors=[query, key, value] + ) + elif self.checkpoint_core_attention and self.training: core_attn_out = self._checkpointed_attention_forward( query, key, diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index 32a755154e9..979a98c2d67 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -497,6 +497,7 @@ def route( hidden_states: torch.Tensor, padding_mask: Optional[torch.Tensor] = None, input_ids: Optional[torch.Tensor] = None, + token_multiplicities: Optional[torch.Tensor] = None, ): """Compute token routing for preprocessing. @@ -506,9 +507,10 @@ def route( """ if padding_mask is not None: padding_mask = padding_mask.transpose(0, 1).bool() - probs, routing_map = apply_module(self.router)( - hidden_states, padding_mask, input_ids=input_ids - ) + router_kwargs = {"input_ids": input_ids} + if token_multiplicities is not None: + router_kwargs["token_multiplicities"] = token_multiplicities + probs, routing_map = apply_module(self.router)(hidden_states, padding_mask, **router_kwargs) return probs, routing_map @maybe_skip_or_early_return_by_cudagraph("preprocess") @@ -717,12 +719,21 @@ def forward( ) self.select_token_dispatcher() + # Keep the tensor in the checkpoint closure: the shared-prefix caller + # removes the scoped attribute before selective recomputation in backward. + token_multiplicities = getattr(self, "_shared_prefix_token_multiplicities", None) + # MoE forward: route -> dispatch -> compute -> combine def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): try: if "route" in self.fwd_execution_map: shared_expert_output = self.shared_experts_compute(hidden_states) - probs, routing_map = self.route(hidden_states, padding_mask, input_ids) + probs, routing_map = self.route( + hidden_states, + padding_mask, + input_ids=input_ids, + token_multiplicities=token_multiplicities, + ) hidden_states, probs = self.preprocess( hidden_states, probs, routing_map, padding_mask ) diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index f89aaf29365..ba5eaea9d8f 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -2,6 +2,9 @@ import functools import math +from collections.abc import Iterator +from contextlib import contextmanager +from contextvars import ContextVar from dataclasses import dataclass from typing import List, Optional, Tuple, Union @@ -1389,6 +1392,59 @@ def apply_biased_logits(logits, std, layer_number=None): return RandomSTEShared.apply(logits, std, layer_number) +_ROUTER_GATING_TOKEN_BLOCK_SIZE: ContextVar[int | None] = ContextVar( + "router_gating_token_block_size", default=None +) + + +@contextmanager +def router_gating_token_blocks(block_size: int = 1024) -> Iterator[None]: + """Keep router GEMM row shapes fixed within a shared-prefix layer execution. + + A short independent star and a larger forest can select different GEMM reduction + algorithms. Even FP32 router differences below 2e-6 can change subsequent BF16 + activations and expert choices. Fixed rows make this reduction independent of + packing. Padding is removed before routing, so it never dispatches extra tokens. + + The caller must enter this scope inside any activation-checkpoint callable so + backward recomputation uses the same router arithmetic as the original forward. + """ + if isinstance(block_size, bool) or not isinstance(block_size, int) or block_size < 1: + raise ValueError("router gating token block size must be a positive integer") + token = _ROUTER_GATING_TOKEN_BLOCK_SIZE.set(block_size) + try: + yield + finally: + _ROUTER_GATING_TOKEN_BLOCK_SIZE.reset(token) + + +def _router_gating_gemm( + inp: torch.Tensor, weight: torch.Tensor, bias: Optional[torch.Tensor], router_dtype: torch.dtype +) -> torch.Tensor: + """Run the existing router GEMM, optionally in fixed, padded row blocks.""" + + def gemm(rows: torch.Tensor) -> torch.Tensor: + if te_general_gemm is not None and router_dtype != torch.float64: + # Preserve the baseline TE bias/output dtype contract. + gemm_bias = bias.to(router_dtype) if bias is not None else None + return te_general_gemm(weight, rows, router_dtype, layout="TN", bias=gemm_bias)[0] + if bias is None: + return torch.mm(rows.to(router_dtype), weight.to(router_dtype).t()) + return torch.addmm( + bias.to(router_dtype), rows.to(router_dtype), weight.to(router_dtype).t() + ) + + block_size = _ROUTER_GATING_TOKEN_BLOCK_SIZE.get() + if block_size is None or inp.shape[0] == 0: + return gemm(inp) + outputs = [] + for start in range(0, inp.shape[0], block_size): + rows = inp[start : start + block_size] + padded = torch.nn.functional.pad(rows, (0, 0, 0, block_size - rows.shape[0])) + outputs.append(gemm(padded)[: rows.shape[0]]) + return torch.cat(outputs, dim=0) + + class RouterGatingLinearFunction(torch.autograd.Function): """ Autograd function for router gating linear. @@ -1421,26 +1477,14 @@ def forward( inp_shape = inp.shape inp = inp.view(-1, inp_shape[-1]) - if te_general_gemm is not None and router_dtype != torch.float64: - # cuBLASLt's non-FP8 bias epilogue expects bias and output to have the same - # dtype. Router parameters may be BF16 while router logits are FP32, so cast the - # small bias vector before passing it to TE. - gemm_bias = bias.to(router_dtype) if bias is not None else None - output = te_general_gemm(weight, inp, router_dtype, layout="TN", bias=gemm_bias)[0] - elif bias is None: - output = torch.mm(inp.to(router_dtype), weight.to(router_dtype).t()) - else: - output = torch.addmm( - bias.to(router_dtype), inp.to(router_dtype), weight.to(router_dtype).t() - ) - + output = _router_gating_gemm(inp, weight, bias, router_dtype) output = output.view(*inp_shape[:-1], -1) return output @staticmethod def backward( ctx, grad_output: torch.Tensor - ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor], None]: + ) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor], None]: """ Backward pass of the RouterGatingLinearFunction function. @@ -1448,7 +1492,7 @@ def backward( grad_output (torch.Tensor): The gradient output. Returns: - Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor], None]: + Tuple[Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor], None]: The gradient input, gradient weight, gradient bias, and None. """ inp, weight, bias = ctx.saved_tensors @@ -1457,21 +1501,33 @@ def backward( inp = inp.view(-1, inp_shape[-1]) grad_output = grad_output.view(-1, grad_shape[-1]) - if te_general_gemm is not None and ctx.router_dtype != torch.float64: - grad_input = te_general_gemm( - weight.to(ctx.router_dtype), grad_output, ctx.router_dtype, layout="NN", grad=True - ) - grad_weight = te_general_gemm( - inp.to(ctx.router_dtype), grad_output, ctx.router_dtype, layout="NT", grad=True - ) - grad_input = grad_input[0].to(ctx.input_dtype) - grad_weight = grad_weight[0].to(ctx.weight_dtype) - else: - grad_input = torch.mm(grad_output, weight.to(ctx.router_dtype)).to(ctx.input_dtype) - grad_weight = torch.mm(grad_output.t(), inp.to(ctx.router_dtype)).to(ctx.weight_dtype) - - grad_bias = grad_output.sum(dim=0).to(ctx.weight_dtype) if bias is not None else None - grad_input = grad_input.view(*inp_shape) + grad_input = grad_weight = grad_bias = None + use_te = te_general_gemm is not None and ctx.router_dtype != torch.float64 + # A frozen router still needs dX to train the hidden states, but has no + # consumer for dW. Autograd's requirements, not router configuration, + # also cover weight-only and bias-only uses of this helper. + if ctx.needs_input_grad[0]: + if use_te: + grad_input = te_general_gemm( + weight.to(ctx.router_dtype), + grad_output, + ctx.router_dtype, + layout="NN", + grad=True, + )[0] + else: + grad_input = torch.mm(grad_output, weight.to(ctx.router_dtype)) + grad_input = grad_input.to(ctx.input_dtype).view(*inp_shape) + if ctx.needs_input_grad[1]: + if use_te: + grad_weight = te_general_gemm( + inp.to(ctx.router_dtype), grad_output, ctx.router_dtype, layout="NT", grad=True + )[0] + else: + grad_weight = torch.mm(grad_output.t(), inp.to(ctx.router_dtype)) + grad_weight = grad_weight.to(ctx.weight_dtype) + if bias is not None and ctx.needs_input_grad[2]: + grad_bias = grad_output.sum(dim=0).to(ctx.weight_dtype) return grad_input, grad_weight, grad_bias, None diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index bc5ff86ae75..d5e7eb2b691 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -43,6 +43,47 @@ _HYBRIDEP_INT16_EXPERT_LIMIT = 1 << 15 +def _expert_bias_token_counts( + routing_map: torch.Tensor, + padding_mask: Optional[torch.Tensor] = None, + token_multiplicities: Optional[torch.Tensor] = None, + *, + num_experts: Optional[int] = None, +) -> torch.Tensor: + """Count logical tokens for boolean maps or dense top-k expert indices. + + Shared prompt rows carry the number of logical copies they replace. Padding + and invalid (-1) dense routes contribute zero, with no data-dependent shapes. + """ + if padding_mask is not None: + padding_mask = padding_mask.reshape(-1) + if padding_mask.shape[0] != routing_map.shape[0]: + raise ValueError("MoE padding mask must contain one value per routed token") + if token_multiplicities is not None: + token_multiplicities = token_multiplicities.reshape(-1) + if token_multiplicities.shape[0] != routing_map.shape[0]: + raise ValueError("MoE token multiplicities must contain one value per routed token") + if token_multiplicities.device != routing_map.device: + raise ValueError("MoE token multiplicities must be on the routing-map device") + else: + token_multiplicities = torch.ones( + routing_map.shape[0], dtype=torch.long, device=routing_map.device + ) + if padding_mask is not None: + token_multiplicities = token_multiplicities.masked_fill(padding_mask, 0) + if routing_map.dtype == torch.bool: + return (routing_map * token_multiplicities.unsqueeze(-1)).sum(dim=0) + if num_experts is None or num_experts < 1: + raise ValueError("Dense expert indices require an explicit positive num_experts") + indices = routing_map.reshape(-1).long() + weights = token_multiplicities.unsqueeze(-1).expand_as(routing_map).reshape(-1) + invalid_routes = indices < 0 + indices = indices.masked_fill(invalid_routes, 0) + weights = weights.masked_fill(invalid_routes, 0) + counts = torch.zeros(num_experts, dtype=token_multiplicities.dtype, device=routing_map.device) + return counts.index_add_(0, indices, weights) + + class Router(ABC, MegatronModule): """Base Router class""" @@ -834,7 +875,10 @@ def apply_input_jitter(self, input: torch.Tensor): @jit_fuser def _apply_expert_bias( - self, routing_map: torch.Tensor, padding_mask: Optional[torch.Tensor] = None + self, + routing_map: torch.Tensor, + padding_mask: Optional[torch.Tensor] = None, + token_multiplicities: Optional[torch.Tensor] = None, ): """ Update expert bias and tokens_per_expert @@ -842,35 +886,46 @@ def _apply_expert_bias( """ if self.enable_expert_bias and torch.is_grad_enabled(): with torch.no_grad(): - use_dense_indices = routing_map.dtype != torch.bool - if padding_mask is not None: - flat_mask = padding_mask.reshape(-1) - assert ( - flat_mask.shape[0] == routing_map.shape[0] - ), f"padding_mask flat {flat_mask.shape} vs routing_map {routing_map.shape}" - if not use_dense_indices: - routing_map = routing_map & (~flat_mask).unsqueeze(-1) - if use_dense_indices: - # Fixed-shape counting: keep every [num_tokens, topk] slot and give padding - # tokens and invalid (-1) routes a zero weight instead of filtering rows, - # which would be a data-dependent shape (nonzero + host sync) inside this - # compiled function and inside the moe_router CUDA graph scope. - expert_indices = routing_map.reshape(-1).to(torch.long) - token_counts = torch.ones_like( - expert_indices, dtype=self.local_tokens_per_expert.dtype + if token_multiplicities is not None: + counts = _expert_bias_token_counts( + routing_map, + padding_mask=padding_mask, + token_multiplicities=token_multiplicities, + num_experts=self.config.num_moe_experts, ) + self.local_tokens_per_expert += counts.to(self.local_tokens_per_expert.dtype) + else: + use_dense_indices = routing_map.dtype != torch.bool if padding_mask is not None: - valid = (~flat_mask).unsqueeze(-1).expand(-1, routing_map.shape[-1]) - token_counts = token_counts * valid.reshape(-1).to(token_counts.dtype) - invalid_routes = expert_indices < 0 - expert_indices = expert_indices.masked_fill(invalid_routes, 0) - token_counts = token_counts.masked_fill(invalid_routes, 0) - if torch.are_deterministic_algorithms_enabled(): - self.local_tokens_per_expert.index_add_(0, expert_indices, token_counts) + flat_mask = padding_mask.reshape(-1) + assert ( + flat_mask.shape[0] == routing_map.shape[0] + ), f"padding_mask flat {flat_mask.shape} vs routing_map {routing_map.shape}" + if not use_dense_indices: + routing_map = routing_map & (~flat_mask).unsqueeze(-1) + if use_dense_indices: + # Fixed-shape counting: keep every [num_tokens, topk] slot and give padding + # tokens and invalid (-1) routes a zero weight instead of filtering rows, + # which would be a data-dependent shape (nonzero + host sync) inside this + # compiled function and inside the moe_router CUDA graph scope. + expert_indices = routing_map.reshape(-1).to(torch.long) + token_counts = torch.ones_like( + expert_indices, dtype=self.local_tokens_per_expert.dtype + ) + if padding_mask is not None: + valid = (~flat_mask).unsqueeze(-1).expand(-1, routing_map.shape[-1]) + token_counts = token_counts * valid.reshape(-1).to(token_counts.dtype) + invalid_routes = expert_indices < 0 + expert_indices = expert_indices.masked_fill(invalid_routes, 0) + token_counts = token_counts.masked_fill(invalid_routes, 0) + if torch.are_deterministic_algorithms_enabled(): + self.local_tokens_per_expert.index_add_(0, expert_indices, token_counts) + else: + self.local_tokens_per_expert.scatter_add_( + 0, expert_indices, token_counts + ) else: - self.local_tokens_per_expert.scatter_add_(0, expert_indices, token_counts) - else: - self.local_tokens_per_expert += routing_map.sum(dim=0) + self.local_tokens_per_expert += routing_map.sum(dim=0) def _hash_routing( self, logits: torch.Tensor, input_ids: torch.Tensor, dense_output: bool = False @@ -942,6 +997,7 @@ def routing( logits: torch.Tensor, padding_mask: Optional[torch.Tensor] = None, input_ids: Optional[torch.Tensor] = None, + token_multiplicities: Optional[torch.Tensor] = None, ): """Top-k routing function @@ -1114,7 +1170,9 @@ def routing( ) # Optionally apply expert bias - self._apply_expert_bias(routing_map, padding_mask=padding_mask) + self._apply_expert_bias( + routing_map, padding_mask=padding_mask, token_multiplicities=token_multiplicities + ) return probs, routing_map @@ -1129,6 +1187,7 @@ def forward( input: torch.Tensor, padding_mask: Optional[torch.Tensor] = None, input_ids: Optional[torch.Tensor] = None, + token_multiplicities: Optional[torch.Tensor] = None, ): """ Forward pass of the router. @@ -1192,7 +1251,12 @@ def forward( logits, self.config.moe_router_force_biased, self.layer_number ) - probs, routing_map = self.routing(logits, padding_mask=padding_mask, input_ids=input_ids) + probs, routing_map = self.routing( + logits, + padding_mask=padding_mask, + input_ids=input_ids, + token_multiplicities=token_multiplicities, + ) return probs, routing_map diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 6b5fe9e48f4..4d107223a54 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -1088,6 +1088,7 @@ def process_mtp_loss( mtp_input_mask: Optional[Tensor] = None, metric_avg_group: Optional[torch.distributed.ProcessGroup] = None, main_hidden_states: Optional[Tensor] = None, + loss_group_lengths: Optional[tuple[int, ...]] = None, ) -> Tensor: """Process Multi-Token Prediction (MTP) loss computation. @@ -1109,6 +1110,9 @@ def process_mtp_loss( packed_seq_params (Optional[PackedSeqParams]): Packed sequence parameters. scale_logits_fn (Optional[Callable[[Tensor], Tensor]]): Optional function to scale logits before loss computation (e.g., MuP output scaling). + loss_group_lengths: Optional CP-local lengths of independently normalized + packed groups. Requires per-token loss and preserves each group's + original/shifted token-count correction when forwards are combined. input_ids (Optional[Tensor]): Input token IDs. Used to derive labels when ``labels`` is None (e.g. RL training), by rolling left to match the SFT label convention (``label[i] = input_id[i + 1]``). Ignored when ``labels`` @@ -1168,6 +1172,21 @@ def process_mtp_loss( # when calculate_per_token_loss is enabled. This ensures MTP gradients are # correctly scaled relative to the main loss gradients in finalize_model_grads. original_num_tokens = loss_mask.sum() + original_group_counts = None + if loss_group_lengths is not None: + if ( + not loss_group_lengths + or any(length < 1 for length in loss_group_lengths) + or sum(loss_group_lengths) != loss_mask.shape[-1] + ): + raise ValueError("MTP loss groups must partition the CP-local packed sequence") + if not config.calculate_per_token_loss: + raise NotImplementedError( + "grouped MTP normalization requires calculate_per_token_loss=True" + ) + original_group_counts = tuple( + part.sum() for part in loss_mask.split(loss_group_lengths, dim=-1) + ) cumulative_mtp_input_mask = None rolled_num_tokens = original_num_tokens @@ -1283,10 +1302,27 @@ def process_mtp_loss( # per-token gradient weighting, we normalize by the rolled token count # and re-scale by the original token count. # Avoid division by zero - num_tokens_safe = torch.clamp(num_tokens, min=1) - mtp_loss_normalized = ( - mtp_loss_scale * mtp_loss * (original_num_tokens / num_tokens_safe) - ) + if loss_group_lengths is None: + num_tokens_safe = torch.clamp(num_tokens, min=1) + mtp_loss_normalized = ( + mtp_loss_scale * mtp_loss * (original_num_tokens / num_tokens_safe) + ) + else: + assert original_group_counts is not None + # Coalescing forwards must not change the pre-existing per-star + # token-count correction, especially at short prompt boundaries. + mtp_loss_normalized = torch.cat( + [ + mtp_loss_scale * part * (original / mask.sum().clamp(min=1)) + for part, mask, original in zip( + mtp_loss.split(loss_group_lengths, dim=-1), + loss_mask.split(loss_group_lengths, dim=-1), + original_group_counts, + strict=True, + ) + ], + dim=-1, + ) hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss_normalized) else: safe_num_tokens = num_tokens.clamp(min=1) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 247c2c50800..a836fa00f82 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -1432,6 +1432,13 @@ class TransformerConfig(ModelParallelConfig): is_hybrid_model: bool = False """ Indicates whether this is a hybrid model. """ + sequence_relative_kernels: bool = False + """Experimental common numerical baseline for dense and shared-prefix execution. + Use sequence-relative Mamba scan chunks and full-context causal attention. + Requires deterministic_tp_reduce_scatter and causal, dropout-free vanilla attention. + Changes rounding relative to the native packed Mamba and TE ring kernels. + """ + mamba_state_dim: int = 128 """The dimensionality of the state representation in Mamba layers.""" @@ -1672,6 +1679,24 @@ def __post_init__(self): "read and write connection." ) + if self.sequence_relative_kernels: + if not self.deterministic_tp_reduce_scatter: + raise ValueError( + "sequence_relative_kernels requires deterministic_tp_reduce_scatter" + ) + if ( + self.attention_dropout != 0 + or self.window_size is not None + or self.qk_clip + or self.log_max_attention_logit + or self.softmax_type != "vanilla" + or self.attn_logit_softcapping is not None + ): + raise ValueError( + "sequence_relative_kernels requires dropout-free vanilla attention without " + "windowing, QK clipping, logit softcapping, or maximum-logit statistics" + ) + # Resolve deprecated attention variant spellings up front so that every consumer # downstream only has to handle the canonical names. Imported lazily because the # spec module imports this one. diff --git a/tests/unit_tests/models/hybrid/test_shared_prefix_port_contracts.py b/tests/unit_tests/models/hybrid/test_shared_prefix_port_contracts.py new file mode 100644 index 00000000000..1fdd3b32db8 --- /dev/null +++ b/tests/unit_tests/models/hybrid/test_shared_prefix_port_contracts.py @@ -0,0 +1,130 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from contextvars import ContextVar, copy_context + +import pytest +import torch + +from megatron.core.models.hybrid.shared_prefix_layout import SharedPrefixLayout +from megatron.core.ssm.mamba_branch_layout import merge_mamba_branches, pack_mamba_branches +from megatron.core.tensor_observation import capture_tensor_observations, observe_tensor +from megatron.core.tensor_parallel.random import _run_recompute_with_observation_suspended +from megatron.core.transformer.moe.moe_utils import router_gating_linear, router_gating_token_blocks +from megatron.core.transformer.moe.router import _expert_bias_token_counts + + +@pytest.mark.parametrize("dense_indices", [False, True]) +@pytest.mark.parametrize("padding", [False, True]) +def test_logical_expert_counts_match_expanded_rows(dense_indices, padding): + routes = torch.tensor([[0, 2], [1, 3], [2, 3], [0, -1]]) + multiplicities = torch.tensor([3, 1, 2, 0]) + padding_mask = torch.tensor([False, True, False, False]) if padding else None + expected = torch.zeros(4, dtype=torch.long) + for row, count in enumerate(multiplicities.tolist()): + if padding and padding_mask[row]: + continue + for expert in routes[row].tolist(): + if expert >= 0: + expected[expert] += count + route_map = routes + if not dense_indices: + route_map = torch.zeros(4, 4, dtype=torch.bool) + for row, experts in enumerate(routes.tolist()): + for expert in experts: + if expert >= 0: + route_map[row, expert] = True + actual = _expert_bias_token_counts(route_map, padding_mask, multiplicities, num_experts=4) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +def test_invalid_expert_count_metadata_rejected(): + routes = torch.tensor([[0, 1], [1, 2]]) + with pytest.raises(ValueError, match="one value per routed token"): + _expert_bias_token_counts(routes, token_multiplicities=torch.ones(3), num_experts=3) + with pytest.raises(ValueError, match="num_experts"): + _expert_bias_token_counts(routes, token_multiplicities=torch.ones(2)) + with pytest.raises(ValueError, match="padding mask"): + _expert_bias_token_counts(routes, padding_mask=torch.zeros(3, dtype=torch.bool)) + + +@pytest.mark.parametrize( + "requires", + [(True, False, False), (False, True, False), (False, False, True), (True, True, True)], +) +@pytest.mark.parametrize("block_size", [None, 3]) +def test_router_frozen_parameters_preserve_requested_gradients(requires, block_size): + generator = torch.Generator().manual_seed(901) + values = [ + torch.randn(shape, generator=generator, dtype=torch.float64) + for shape in [(7, 1, 5), (4, 5), (4,)] + ] + actual_inputs = [value.clone().requires_grad_(need) for value, need in zip(values, requires)] + reference_inputs = [value.clone().requires_grad_(need) for value, need in zip(values, requires)] + if block_size is None: + actual = router_gating_linear(*actual_inputs, torch.float64) + else: + with router_gating_token_blocks(block_size): + actual = router_gating_linear(*actual_inputs, torch.float64) + reference = torch.nn.functional.linear(*reference_inputs) + cotangent = torch.randn(actual.shape, generator=generator, dtype=torch.float64) + actual.backward(cotangent) + reference.backward(cotangent) + torch.testing.assert_close(actual, reference, rtol=1e-12, atol=1e-12) + for value, reference_value, need in zip(actual_inputs, reference_inputs, requires): + if need: + torch.testing.assert_close(value.grad, reference_value.grad, rtol=1e-12, atol=1e-12) + else: + assert value.grad is None + + +@pytest.mark.parametrize("tail", [0, 2, 3]) +def test_mamba_branch_layout_gradients_match_slice_reference(tail): + lengths = (2, 0, 5) + values = torch.randn(10, 1, 2, dtype=torch.float64, requires_grad=True) + reference = values.detach().clone().requires_grad_() + actual = pack_mamba_branches(values, prefix_len=3, tail_len=tail, completion_lens=lengths) + branches = reference.new_zeros(tail + 5, 3, 2) + offset = 3 + for branch, length in enumerate(lengths): + branches[:tail, branch] = reference[3 - tail : 3, 0] + branches[tail : tail + length, branch] = reference[offset : offset + length, 0] + offset += length + cotangent = torch.randn_like(actual) + actual.backward(cotangent) + branches.backward(cotangent) + torch.testing.assert_close(actual, branches, rtol=0, atol=0) + torch.testing.assert_close(values.grad, reference.grad, rtol=0, atol=0) + head = torch.randn(3 - tail, 1, 2, dtype=torch.float64, requires_grad=True) + branch_input = actual.detach().requires_grad_() + assert torch.autograd.gradcheck( + lambda h, b: merge_mamba_branches(h, b, tail_len=tail, completion_lens=lengths), + (head, branch_input), + ) + + +def test_layout_prompt_multiplicities_equal_dense_gather_counts(): + layout = SharedPrefixLayout(3, (2, 5, 1)) + indices = torch.cat(layout.dense_branch_indices("cpu")) + expected = torch.bincount(indices, minlength=layout.total_len).float() + torch.testing.assert_close( + layout.padded_token_multiplicities(layout.total_len, "cpu"), expected, rtol=0, atol=0 + ) + assert layout.position_ids("cpu").tolist() == [0, 1, 2, 3, 4, 3, 4, 5, 6, 7, 3] + + +def test_recompute_restores_context_without_duplicate_observations(): + observed = [] + shared_scope = ContextVar("test_shared_scope", default=None) + token = shared_scope.set("shared-forward") + with capture_tensor_observations(lambda *args: observed.append(args), frozenset({"test"})): + saved_context = copy_context() + shared_scope.reset(token) + + def recompute(): + observe_tensor(None, "test", "test", torch.ones(1)) + return shared_scope.get() + + value = saved_context.copy().run(_run_recompute_with_observation_suspended, recompute) + assert value == "shared-forward" + assert observed == [] + assert shared_scope.get() is None