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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
92 changes: 92 additions & 0 deletions docs/developer/shared_prefix.md
Original file line number Diff line number Diff line change
@@ -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.
66 changes: 65 additions & 1 deletion megatron/core/extensions/transformer_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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")
Expand Down
6 changes: 6 additions & 0 deletions megatron/core/model_parallel_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
Loading