diff --git a/gpt_builders.py b/gpt_builders.py index 33af72ecfef..1ace7267de6 100644 --- a/gpt_builders.py +++ b/gpt_builders.py @@ -2,6 +2,7 @@ from megatron.core.models.gpt import GPTModel from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_stage_input_cp_partition_mode, get_transformer_block_with_experimental_attention_variant_spec, get_transformer_layer_with_experimental_attention_variant_spec, ) @@ -17,6 +18,7 @@ get_gpt_heterogeneous_layer_spec, ) from megatron.core.transformer.spec_utils import import_module +from megatron.core.utils import get_pg_rank from megatron.training import get_args, print_rank_0 from megatron.training.arguments import core_transformer_config_from_args from megatron.training.yaml_arguments import core_transformer_config_from_yaml @@ -29,14 +31,28 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_ config = core_transformer_config_from_yaml(args, "language_model") else: config = core_transformer_config_from_args(args) + cp_stage_entry_partition_mode = ( + "zigzag" if config.cp_partition_mode == "auto" else config.cp_partition_mode + ) if args.spec is not None: transformer_layer_spec = import_module(args.spec) else: use_te = args.transformer_impl == "transformer_engine" if args.experimental_attention_variant is not None: + pp_rank = ( + get_pg_rank(pg_collection.pp) + if pg_collection is not None and hasattr(pg_collection, "pp") + else None + ) + if config.cp_partition_mode == "auto": + cp_stage_entry_partition_mode = ( + get_experimental_attention_variant_stage_input_cp_partition_mode( + config=config, vp_stage=vp_stage, pp_rank=pp_rank + ) + ) transformer_layer_spec = get_transformer_block_with_experimental_attention_variant_spec( - config=config, vp_stage=vp_stage + config=config, vp_stage=vp_stage, pp_rank=pp_rank ) elif args.num_experts: # Define the decoder block spec @@ -107,6 +123,7 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_ mtp_block_spec=mtp_block_spec, vp_stage=vp_stage, pg_collection=pg_collection, + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, ) return model diff --git a/hybrid_builders.py b/hybrid_builders.py index d95002e1a21..64a482caa77 100644 --- a/hybrid_builders.py +++ b/hybrid_builders.py @@ -1,6 +1,9 @@ # Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_inference_stack_spec +from megatron.core.models.hybrid.hybrid_layer_allocation import ( + get_hybrid_stage_input_cp_partition_mode_for_stage, +) from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.transformer import MLATransformerConfig, TransformerConfig from megatron.core.transformer.spec_utils import ModuleSpec, import_module @@ -75,6 +78,17 @@ def hybrid_builder(args, pre_process, post_process, vp_stage=None, config=None, else: raise ValueError("You must provide a valid hybrid layer spec via --spec") + cp_stage_entry_partition_mode = config.cp_partition_mode + if config.cp_partition_mode == "auto": + cp_stage_entry_partition_mode = get_hybrid_stage_input_cp_partition_mode_for_stage( + config, + args.hybrid_layer_pattern, + getattr(pg_collection, "pp", None), + vp_stage, + first_stage_layers=config.num_layers_in_first_pipeline_stage, + last_stage_layers=config.num_layers_in_last_pipeline_stage, + ) + model = HybridModel( config=config, hybrid_stack_spec=hybrid_stack_spec, @@ -91,6 +105,7 @@ def hybrid_builder(args, pre_process, post_process, vp_stage=None, config=None, rotary_base=args.rotary_base, pg_collection=pg_collection, vp_stage=vp_stage, + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, ) for l in range(model.decoder.num_layers_per_pipeline_rank): diff --git a/megatron/core/context_parallel_layout.py b/megatron/core/context_parallel_layout.py deleted file mode 100644 index 44014581fd5..00000000000 --- a/megatron/core/context_parallel_layout.py +++ /dev/null @@ -1,307 +0,0 @@ -# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. - -"""Context parallel tensor layout helpers.""" - -from typing import List, Optional, Tuple - -import torch - -from megatron.core.tensor_parallel import all_to_all - - -def get_thd_context_parallel_rank_indices( - cu_seqlens: torch.Tensor, cp_size: int, cp_rank: int, layout: str -) -> torch.Tensor: - """Return global THD token indices owned by one CP rank in a layout. - - Args: - cu_seqlens: Global packed-sequence cumulative lengths before CP partitioning. - cp_size: Context-parallel group size. - cp_rank: Context-parallel rank. - layout: Either ``"zigzag"`` or ``"contiguous"``. - - The returned indices are ordered exactly as the rank-local THD tensor is stored. - ``"zigzag"`` follows Megatron's per-sequence load-balanced chunk order; ``"contiguous"`` - partitions the flattened packed THD buffer into rank-contiguous spans. - """ - if layout not in ("zigzag", "contiguous"): - raise ValueError(f"Unsupported context-parallel layout {layout!r}.") - if cp_size < 1: - raise ValueError(f"cp_size must be >= 1, got {cp_size}.") - if not 0 <= cp_rank < cp_size: - raise ValueError(f"cp_rank must be in [0, {cp_size}), got {cp_rank}.") - if cu_seqlens.dim() != 1: - raise ValueError(f"cu_seqlens must be 1-D, got shape {tuple(cu_seqlens.shape)}.") - - cu = cu_seqlens.to(dtype=torch.long) - if cu.numel() == 0 or cu[0].item() != 0: - raise ValueError(f"cu_seqlens must start at 0, got {cu_seqlens}.") - - if torch.any(torch.diff(cu) < 0): - raise ValueError(f"cu_seqlens must be nondecreasing, got {cu_seqlens}.") - - nonduplicate_boundaries = torch.ones(cu.numel(), device=cu.device, dtype=torch.bool) - nonduplicate_boundaries[1:] = cu[1:] != cu[:-1] - cu = cu[nonduplicate_boundaries] - - total_tokens = int(cu[-1].item()) - positions = torch.arange(total_tokens, device=cu.device, dtype=torch.long) - if total_tokens == 0: - return positions - - seq_lens = torch.diff(cu) - chunk_divisor = 2 * cp_size - if torch.any(seq_lens % chunk_divisor != 0): - raise ValueError( - "All packed sequence lengths must be divisible by " - f"2 * cp_size ({chunk_divisor}) for zigzag/contiguous CP layout conversion, " - f"got {seq_lens}." - ) - - if layout == "contiguous": - part_len = total_tokens // cp_size - rank_start = cp_rank * part_len - return positions[rank_start : rank_start + part_len] - - seq_idx = torch.bucketize(positions, cu[1:], right=True) - global_starts = cu[:-1] - pos_in_seq = positions - global_starts[seq_idx] - chunk_lens = (seq_lens // chunk_divisor)[seq_idx] - chunk = pos_in_seq // chunk_lens - offset = pos_in_seq - chunk * chunk_lens - - owner = torch.where(chunk < cp_size, chunk, 2 * cp_size - chunk - 1) - local_slot = torch.where(chunk < cp_size, torch.zeros_like(chunk), torch.ones_like(chunk)) - - local_starts = (global_starts // cp_size)[seq_idx] - local_pos = local_starts + local_slot * chunk_lens + offset - - rank_mask = owner == cp_rank - rank_positions = positions[rank_mask] - rank_local_pos = local_pos[rank_mask] - return rank_positions[torch.argsort(rank_local_pos)] - - -def zigzag_to_contiguous_chunks( - x: torch.Tensor, - cp_group: torch.distributed.ProcessGroup, - seq_dim: int = 0, - cu_seqlens: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Permute CP chunks from Megatron zigzag layout to contiguous-time layout. - - SBHD tensors have two equal chunks per rank along ``seq_dim`` and use a - chunk-level all-to-all. THD tensors pass global ``cu_seqlens`` and use one - packed-token all-to-all over the whole local THD tensor. - """ - if cu_seqlens is not None: - return _zigzag_contiguous_thd_swap( - x, cp_group, seq_dim, cu_seqlens, source_layout="zigzag", target_layout="contiguous" - ) - return _zigzag_contiguous_chunk_swap(x, cp_group, seq_dim, to_contiguous=True) - - -def contiguous_to_zigzag_chunks( - x: torch.Tensor, - cp_group: torch.distributed.ProcessGroup, - seq_dim: int = 0, - cu_seqlens: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Inverse of :func:`zigzag_to_contiguous_chunks`.""" - if cu_seqlens is not None: - return _zigzag_contiguous_thd_swap( - x, cp_group, seq_dim, cu_seqlens, source_layout="contiguous", target_layout="zigzag" - ) - return _zigzag_contiguous_chunk_swap(x, cp_group, seq_dim, to_contiguous=False) - - -def _zigzag_contiguous_thd_swap( - x: torch.Tensor, - cp_group: Optional[torch.distributed.ProcessGroup], - seq_dim: int, - cu_seqlens: torch.Tensor, - source_layout: str, - target_layout: str, -) -> torch.Tensor: - """Single-all-to-all THD permutation between zigzag and contiguous layouts. - - The packed THD tensor stays packed: we first group local tokens by their - target CP rank, exchange those groups once, then scatter received tokens - back into the target rank-local order. - """ - cp_size = cp_group.size() if cp_group is not None else 1 - if cp_size == 1: - return x - cp_rank = cp_group.rank() - - if seq_dim != 0: - x = x.movedim(seq_dim, 0) - x = x.contiguous() - - cu = cu_seqlens.to(device=x.device, dtype=torch.long) - # TODO: Let a future CP layout scheduler precompute this routing once per - # microbatch from immutable cu_seqlens and pass it through both THD swaps. - # Do not cache it across microbatches because packed sequence boundaries change. - source_by_rank = [ - get_thd_context_parallel_rank_indices(cu, cp_size, rank, source_layout) - for rank in range(cp_size) - ] - target_by_rank = [ - get_thd_context_parallel_rank_indices(cu, cp_size, rank, target_layout) - for rank in range(cp_size) - ] - - local_source_indices = source_by_rank[cp_rank] - local_target_indices = target_by_rank[cp_rank] - if x.size(0) != local_source_indices.numel(): - raise ValueError( - f"Local THD tensor length ({x.size(0)}) does not match {source_layout} " - f"rank-{cp_rank} partition length ({local_source_indices.numel()})." - ) - - total_tokens = int(cu[-1].item()) - target_owner = torch.empty(total_tokens, device=x.device, dtype=torch.long) - target_local_pos = torch.empty(total_tokens, device=x.device, dtype=torch.long) - for rank, indices in enumerate(target_by_rank): - target_owner[indices] = rank - target_local_pos[indices] = torch.arange(indices.numel(), device=x.device) - - local_target_owner = target_owner[local_source_indices] - local_target_pos = target_local_pos[local_source_indices] - - send_parts: List[torch.Tensor] = [] - input_split_sizes: List[int] = [] - for dst_rank in range(cp_size): - dst_mask = local_target_owner == dst_rank - dst_rows = dst_mask.nonzero(as_tuple=False).flatten() - if dst_rows.numel() > 0: - dst_rows = dst_rows[torch.argsort(local_target_pos[dst_rows])] - send_part = x.index_select(0, dst_rows) - else: - send_part = x.narrow(0, 0, 0) - send_parts.append(send_part) - input_split_sizes.append(send_part.size(0)) - send_buf = torch.cat(send_parts, dim=0).contiguous() - - output_split_sizes: List[int] = [] - recv_target_positions: List[torch.Tensor] = [] - for src_rank in range(cp_size): - src_indices = source_by_rank[src_rank] - src_to_this_rank = target_owner[src_indices] == cp_rank - recv_global_indices = src_indices[src_to_this_rank] - if recv_global_indices.numel() > 0: - recv_positions = target_local_pos[recv_global_indices] - recv_positions = recv_positions[torch.argsort(recv_positions)] - else: - recv_positions = local_target_indices.narrow(0, 0, 0) - recv_target_positions.append(recv_positions) - output_split_sizes.append(recv_positions.numel()) - - recv_buf = all_to_all(cp_group, send_buf, output_split_sizes, input_split_sizes) - - out_shape = (local_target_indices.numel(),) + tuple(x.shape[1:]) - out = x.new_empty(out_shape) - offset = 0 - for recv_positions in recv_target_positions: - recv_len = recv_positions.numel() - if recv_len > 0: - out[recv_positions] = recv_buf[offset : offset + recv_len] - offset += recv_len - - if seq_dim != 0: - out = out.movedim(0, seq_dim) - return out.contiguous() - - -def _zigzag_contiguous_chunk_swap( - x: torch.Tensor, - cp_group: Optional[torch.distributed.ProcessGroup], - seq_dim: int, - to_contiguous: bool, -) -> torch.Tensor: - """Single-all-to-all chunk permutation between zigzag and contiguous layouts. - - Each rank holds exactly two chunks along ``seq_dim``. The mapping from - local (rank, slot) to (rank, slot) in the target layout is deterministic - and depends only on ``cp_size`` and ``cp_rank``, so we pack send data in - destination-rank order and use one ``all_to_all_single`` with unequal - splits to route each chunk to its target rank. - """ - cp_size = cp_group.size() if cp_group is not None else 1 - if cp_size == 1: - return x - cp_rank = cp_group.rank() - - # Work with seq_dim at position 0. - if seq_dim != 0: - x = x.movedim(seq_dim, 0) - x = x.contiguous() - - seq_len_local = x.size(0) - assert seq_len_local % 2 == 0, ( - f"zigzag/contiguous chunk swap requires an even local sequence length, " - f"got {seq_len_local}." - ) - chunk_len = seq_len_local // 2 - - def _rank_to_chunks(rank: int, in_zigzag: bool) -> Tuple[int, int]: - """Global chunk indices at (slot 0, slot 1) for this rank.""" - if in_zigzag: - return (rank, 2 * cp_size - rank - 1) - return (2 * rank, 2 * rank + 1) - - def _chunk_to_dest(chunk_idx: int, target_zigzag: bool) -> Tuple[int, int]: - """Destination (rank, slot) for a given global chunk index in the target layout.""" - if target_zigzag: - if chunk_idx < cp_size: - return chunk_idx, 0 - return 2 * cp_size - chunk_idx - 1, 1 - return chunk_idx // 2, chunk_idx % 2 - - source_in_zigzag = to_contiguous - target_in_zigzag = not to_contiguous - - local_chunk_indices = _rank_to_chunks(cp_rank, source_in_zigzag) - local_dests = [_chunk_to_dest(c, target_in_zigzag) for c in local_chunk_indices] - - # Pack the send buffer so chunks are ordered by (dst_rank, dst_slot). - local_slot_order = sorted(range(2), key=lambda s: local_dests[s]) - local_chunks = [x[:chunk_len], x[chunk_len:]] - send_buf = torch.cat([local_chunks[s] for s in local_slot_order], dim=0).contiguous() - - input_split_chunks = [0] * cp_size - for dst_rank, _ in local_dests: - input_split_chunks[dst_rank] += 1 - - # Mirror every source rank's packing logic so we know which received chunk - # belongs in which local target slot. - output_split_chunks = [0] * cp_size - recv_dst_slots_per_source: List[List[int]] = [[] for _ in range(cp_size)] - for src in range(cp_size): - src_chunks = _rank_to_chunks(src, source_in_zigzag) - src_dests = [_chunk_to_dest(c, target_in_zigzag) for c in src_chunks] - src_slot_order = sorted(range(2), key=lambda s: src_dests[s]) - for s in src_slot_order: - dst_rank, dst_slot = src_dests[s] - if dst_rank == cp_rank: - output_split_chunks[src] += 1 - recv_dst_slots_per_source[src].append(dst_slot) - - input_split_sizes = [n * chunk_len for n in input_split_chunks] - output_split_sizes = [n * chunk_len for n in output_split_chunks] - - recv_buf = all_to_all(cp_group, send_buf, output_split_sizes, input_split_sizes) - - # Reassemble local chunks in target-layout slot order. - target_slots: List[Optional[torch.Tensor]] = [None, None] - offset = 0 - for src in range(cp_size): - for dst_slot in recv_dst_slots_per_source[src]: - target_slots[dst_slot] = recv_buf[offset : offset + chunk_len] - offset += chunk_len - assert all(t is not None for t in target_slots), "Incomplete chunk reassembly in CP swap" - - out = torch.cat(target_slots, dim=0) - if seq_dim != 0: - out = out.movedim(0, seq_dim) - return out.contiguous() diff --git a/megatron/core/context_parallel_layout/__init__.py b/megatron/core/context_parallel_layout/__init__.py new file mode 100644 index 00000000000..b0820c00d04 --- /dev/null +++ b/megatron/core/context_parallel_layout/__init__.py @@ -0,0 +1,61 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Context parallel sequence partition-mode helpers. + +This package preserves the historical ``megatron.core.context_parallel_layout`` +import surface while splitting the implementation by responsibility. + +Ownership summary: + +- model builders choose the pipeline-stage input CP layout; +- blocks convert rank-local sequence tensors between layer preferences; +- model postprocess restores the public output boundary to the input layout; +- MTP validates its inner-layer layout preference but does not own outer conversion. +""" + +from typing import Literal + +CpPartitionMode = Literal["zigzag", "contiguous"] + +from megatron.core.context_parallel_layout.conversion import ( + CpPartitionModeConverter, + contiguous_to_zigzag_chunks, + convert_cp_partition_mode, + convert_module_input_tensors_cp_partition_mode, + zigzag_to_contiguous_chunks, +) +from megatron.core.context_parallel_layout.metadata import ( + get_packed_seq_params_cp_partition_cu_seqlens, + replace_packed_seq_params_cp_partition_mode, +) +from megatron.core.context_parallel_layout.policy import ( + get_context_parallel_layout_chunk_indices, + get_preferred_cp_partition_mode_for_layer, + get_stage_entry_partition_mode, +) +from megatron.core.context_parallel_layout.routes import ( + build_thd_cp_partition_route, + decode_thd_cp_partition_route, + get_thd_context_parallel_rank_indices, + get_thd_cp_partition_route, + prebuild_thd_cp_partition_routes, +) + +__all__ = [ + "CpPartitionMode", + "CpPartitionModeConverter", + "build_thd_cp_partition_route", + "contiguous_to_zigzag_chunks", + "convert_cp_partition_mode", + "convert_module_input_tensors_cp_partition_mode", + "decode_thd_cp_partition_route", + "get_context_parallel_layout_chunk_indices", + "get_packed_seq_params_cp_partition_cu_seqlens", + "get_preferred_cp_partition_mode_for_layer", + "get_stage_entry_partition_mode", + "get_thd_cp_partition_route", + "get_thd_context_parallel_rank_indices", + "prebuild_thd_cp_partition_routes", + "replace_packed_seq_params_cp_partition_mode", + "zigzag_to_contiguous_chunks", +] diff --git a/megatron/core/context_parallel_layout/conversion.py b/megatron/core/context_parallel_layout/conversion.py new file mode 100644 index 00000000000..bf63886dd45 --- /dev/null +++ b/megatron/core/context_parallel_layout/conversion.py @@ -0,0 +1,564 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Tensor operations for converting between CP partition modes.""" + +import copy +from typing import Any, Callable, List, Optional, Tuple, Union + +import torch + +from megatron.core.context_parallel_layout.metadata import ( + get_packed_seq_params_cp_partition_cu_seqlens, +) +from megatron.core.context_parallel_layout.routes import ( + _cp_layout_nvtx_range, + build_thd_cp_partition_route, + decode_thd_cp_partition_route, + get_thd_cp_partition_route, +) +from megatron.core.context_parallel_layout import CpPartitionMode + + +class CpPartitionModeConverter: + """Convert tensors across one CP layout edge.""" + + def __init__( + self, + *, + cp_group: Optional[torch.distributed.ProcessGroup], + packed_seq_params: Optional[Any], + source_partition_mode: Optional[CpPartitionMode], + target_partition_mode: Optional[CpPartitionMode], + config: Any, + tp_group: Optional[torch.distributed.ProcessGroup] = None, + ) -> None: + self.cp_group = cp_group + self.packed_seq_params = packed_seq_params + self.source_partition_mode = source_partition_mode + self.target_partition_mode = target_partition_mode + self.config = config + self.tp_group = tp_group + if ( + self.conversion_needed + and getattr(self.packed_seq_params, "qkv_format", None) == "thd" + and self.config.cuda_graph_impl == "full_iteration" + ): + raise ValueError( + "Full-iteration CUDA graph is not supported for THD CP layout conversion: " + f"source={self.source_partition_mode!r}, target={self.target_partition_mode!r}." + ) + + @property + def conversion_needed(self) -> bool: + """Return whether this edge needs a real layout conversion.""" + return ( + self.source_partition_mode != self.target_partition_mode + and self.cp_group is not None + and self.cp_group.size() > 1 + ) + + def assert_no_dense_attention_inputs( + self, + *, + attention_mask: Optional[torch.Tensor] = None, + attention_bias: Optional[torch.Tensor] = None, + hidden_states: Optional[torch.Tensor] = None, + ) -> None: + """Reject dense attention tensors when this edge would reorder tokens.""" + if not self.conversion_needed: + return + if attention_mask is not None: + self._raise_unsupported_dense_attention( + "an explicit attention_mask", hidden_states=hidden_states + ) + if attention_bias is not None: + self._raise_unsupported_dense_attention( + "attention_bias", hidden_states=hidden_states + ) + + def convert( + self, + value: Any, + *, + seq_dim: Union[int, Callable[[torch.Tensor], int]] = 0, + sequence_parallel: bool = False, + ) -> Any: + """Convert a tensor or nested tensor container across this layout edge.""" + if not self.conversion_needed or value is None: + return value + # Nested values may contain optional tensors; traverse containers while + # preserving their original shape. + if isinstance(value, tuple): + return tuple( + self.convert(part, seq_dim=seq_dim, sequence_parallel=sequence_parallel) + for part in value + ) + if isinstance(value, list): + return [ + self.convert(part, seq_dim=seq_dim, sequence_parallel=sequence_parallel) + for part in value + ] + if not torch.is_tensor(value): + return value + + resolved_seq_dim = seq_dim(value) if callable(seq_dim) else seq_dim + return convert_cp_partition_mode( + value, + self.cp_group, + source_partition_mode=self.source_partition_mode, + target_partition_mode=self.target_partition_mode, + seq_dim=resolved_seq_dim, + cu_seqlens=get_packed_seq_params_cp_partition_cu_seqlens(self.packed_seq_params), + sequence_parallel=sequence_parallel, + tp_group=self.tp_group, + thd_cp_partition_route=get_thd_cp_partition_route( + self.packed_seq_params, + self.source_partition_mode, + self.target_partition_mode, + ), + ) + + def _raise_unsupported_dense_attention( + self, + tensor_name: str, + *, + hidden_states: Optional[torch.Tensor], + ) -> None: + hidden_shape = tuple(hidden_states.shape) if hidden_states is not None else None + raise NotImplementedError( + "Changing CP partition mode with " + f"{tensor_name} is not supported yet: " + f"source={self.source_partition_mode!r}, " + f"target={self.target_partition_mode!r}, " + f"qkv_format={getattr(self.packed_seq_params, 'qkv_format', None)!r}, " + f"hidden_shape={hidden_shape}." + ) + + +def convert_module_input_tensors_cp_partition_mode( + *, + hidden_states: torch.Tensor, + packed_seq_params: Optional[Any], + cp_group: Optional[torch.distributed.ProcessGroup], + tp_group: Optional[torch.distributed.ProcessGroup], + target_partition_mode: CpPartitionMode, + sequence_parallel: bool, + config: Any, + attention_mask: Optional[torch.Tensor] = None, + attention_bias: Optional[torch.Tensor] = None, + key_value_states: Optional[torch.Tensor] = None, +) -> Tuple[ + torch.Tensor, + Optional[Any], + Optional[CpPartitionModeConverter], +]: + """Convert a module's rank-local sequence tensors to a target CP layout. + + This helper performs the common "entry conversion" pattern used by modules + that need to consume a different CP layout than their caller supplied. It + returns a converter for the opposite edge so the module output can be + converted back to the original input layout. + """ + if cp_group is None or cp_group.size() <= 1: + return hidden_states, packed_seq_params, None + + source_partition_mode = getattr(packed_seq_params, "cp_partition_mode", None) + if source_partition_mode is None: + raise ValueError( + "PackedSeqParams.cp_partition_mode is required before module input CP layout " + "conversion when context parallelism is active." + ) + if source_partition_mode == target_partition_mode: + return hidden_states, packed_seq_params, None + + input_to_target_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=source_partition_mode, + target_partition_mode=target_partition_mode, + config=config, + tp_group=tp_group, + ) + input_to_target_converter.assert_no_dense_attention_inputs( + attention_mask=attention_mask, + attention_bias=attention_bias, + hidden_states=hidden_states, + ) + if key_value_states is not None: + raise NotImplementedError( + "Changing CP partition mode with cross-attention key/value states is not supported " + f"yet: source={source_partition_mode!r}, target={target_partition_mode!r}." + ) + hidden_states = input_to_target_converter.convert( + hidden_states, + seq_dim=0, + sequence_parallel=sequence_parallel, + ) + + local_packed_seq_params = packed_seq_params + if packed_seq_params is not None: + local_packed_seq_params = copy.copy(packed_seq_params) + local_packed_seq_params.cp_partition_mode = target_partition_mode + target_to_input_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=local_packed_seq_params, + source_partition_mode=target_partition_mode, + target_partition_mode=source_partition_mode, + config=config, + tp_group=tp_group, + ) + return ( + hidden_states, + local_packed_seq_params, + target_to_input_converter, + ) + + +def zigzag_to_contiguous_chunks( + x: torch.Tensor, + cp_group: torch.distributed.ProcessGroup, + seq_dim: int = 0, + cu_seqlens: Optional[torch.Tensor] = None, + thd_cp_partition_route: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Permute CP chunks from Megatron zigzag layout to contiguous-time layout. + + SBHD tensors have two equal chunks per rank along ``seq_dim`` and use a + chunk-level all-to-all. THD tensors pass global ``cu_seqlens`` and use one + packed-token all-to-all over the whole local THD tensor. + """ + if cu_seqlens is not None: + return _zigzag_contiguous_thd_swap( + x, + cp_group, + seq_dim, + cu_seqlens, + source_partition_mode="zigzag", + target_partition_mode="contiguous", + thd_cp_partition_route=thd_cp_partition_route, + ) + return _zigzag_contiguous_chunk_swap(x, cp_group, seq_dim, to_contiguous=True) + + +def contiguous_to_zigzag_chunks( + x: torch.Tensor, + cp_group: torch.distributed.ProcessGroup, + seq_dim: int = 0, + cu_seqlens: Optional[torch.Tensor] = None, + thd_cp_partition_route: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Inverse of :func:`zigzag_to_contiguous_chunks`.""" + if cu_seqlens is not None: + return _zigzag_contiguous_thd_swap( + x, + cp_group, + seq_dim, + cu_seqlens, + source_partition_mode="contiguous", + target_partition_mode="zigzag", + thd_cp_partition_route=thd_cp_partition_route, + ) + return _zigzag_contiguous_chunk_swap(x, cp_group, seq_dim, to_contiguous=False) + + +def convert_cp_partition_mode( + x: torch.Tensor, + cp_group: Optional[torch.distributed.ProcessGroup], + *, + source_partition_mode: Optional[str], + target_partition_mode: Optional[str], + seq_dim: int = 0, + cu_seqlens: Optional[torch.Tensor] = None, + sequence_parallel: bool = False, + tp_group: Optional[torch.distributed.ProcessGroup] = None, + thd_cp_partition_route: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Convert a sequence tensor between CP zigzag and contiguous layouts. + + With sequence parallel enabled, the baseline path gathers the full CP-local + sequence on each TP rank, performs the CP layout conversion, then scatters + back to the original SP sharding. + """ + # TODO(yuzhongw): implement a direct TPxCP layout conversion path if the + # gather-convert-scatter fallback becomes a bottleneck. + + if source_partition_mode == target_partition_mode: + return x + + cp_size = cp_group.size() if cp_group is not None else 1 + if cp_size == 1: + return x + + if source_partition_mode not in ("zigzag", "contiguous") or target_partition_mode not in ( + "zigzag", + "contiguous", + ): + cp_rank = cp_group.rank() if cp_group is not None else 0 + raise ValueError( + f"Unsupported CP partition mode conversion " + f"{source_partition_mode!r} -> {target_partition_mode!r}; " + f"shape={tuple(x.shape)}, seq_dim={seq_dim}, cp_size={cp_size}, cp_rank={cp_rank}." + ) + + if sequence_parallel and tp_group is not None and tp_group.size() > 1: + from megatron.core.tensor_parallel.mappings import ( + gather_from_sequence_parallel_region, + scatter_to_sequence_parallel_region, + ) + + moved = x.movedim(seq_dim, 0) if seq_dim != 0 else x + # This gather is only used to run a duplicated CP layout permutation before + # scattering back to SP shards. Its backward must split, not reduce-scatter; + # otherwise every TP rank contributes the same full-sequence gradient. + gathered = gather_from_sequence_parallel_region( + moved, + tensor_parallel_output_grad=False, + group=tp_group, + ) + converted = _convert_cp_partition_mode_full_sequence( + gathered, + cp_group, + source_partition_mode=source_partition_mode, + target_partition_mode=target_partition_mode, + seq_dim=0, + cu_seqlens=cu_seqlens, + thd_cp_partition_route=thd_cp_partition_route, + ) + scattered = scatter_to_sequence_parallel_region(converted, group=tp_group) + return scattered.movedim(0, seq_dim).contiguous() if seq_dim != 0 else scattered + + return _convert_cp_partition_mode_full_sequence( + x, + cp_group, + source_partition_mode=source_partition_mode, + target_partition_mode=target_partition_mode, + seq_dim=seq_dim, + cu_seqlens=cu_seqlens, + thd_cp_partition_route=thd_cp_partition_route, + ) + + +def _convert_cp_partition_mode_full_sequence( + x: torch.Tensor, + cp_group: Optional[torch.distributed.ProcessGroup], + *, + source_partition_mode: CpPartitionMode, + target_partition_mode: CpPartitionMode, + seq_dim: int, + cu_seqlens: Optional[torch.Tensor], + thd_cp_partition_route: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Convert a tensor whose sequence dim contains the full CP-local sequence.""" + if source_partition_mode == "zigzag" and target_partition_mode == "contiguous": + return zigzag_to_contiguous_chunks( + x, + cp_group, + seq_dim=seq_dim, + cu_seqlens=cu_seqlens, + thd_cp_partition_route=thd_cp_partition_route, + ) + if source_partition_mode == "contiguous" and target_partition_mode == "zigzag": + return contiguous_to_zigzag_chunks( + x, + cp_group, + seq_dim=seq_dim, + cu_seqlens=cu_seqlens, + thd_cp_partition_route=thd_cp_partition_route, + ) + raise ValueError( + f"Unsupported CP partition mode conversion " + f"{source_partition_mode!r} -> {target_partition_mode!r}; " + f"shape={tuple(x.shape)}, seq_dim={seq_dim}." + ) + + +def _pack_thd_cp_route_send_buffer( + x: torch.Tensor, local_source_length: int, send_rows: Optional[torch.Tensor] +) -> torch.Tensor: + if local_source_length == 0: + return x.narrow(0, 0, 0) + if send_rows is None: + return x + return x.index_select(0, send_rows) + + +def _scatter_thd_cp_route_recv_buffer( + recv_buf: torch.Tensor, recv_rows: Optional[torch.Tensor], out_shape: Tuple[int, ...] +) -> torch.Tensor: + if recv_rows is None: + return recv_buf + out = recv_buf.new_empty(out_shape) + if recv_rows.numel() > 0: + out.index_copy_(0, recv_rows, recv_buf) + return out + + +def _zigzag_contiguous_thd_swap( + x: torch.Tensor, + cp_group: Optional[torch.distributed.ProcessGroup], + seq_dim: int, + cu_seqlens: torch.Tensor, + source_partition_mode: str, + target_partition_mode: str, + thd_cp_partition_route: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Single-all-to-all THD permutation between zigzag and contiguous layouts. + + The packed THD tensor stays packed: we first group local tokens by their + target CP rank, exchange those groups once, then scatter received tokens + back into the target rank-local order. + """ + cp_size = cp_group.size() if cp_group is not None else 1 + if cp_size == 1: + return x + cp_rank = cp_group.rank() + from megatron.core.tensor_parallel.mappings import all_to_all + + conversion_name = f"{source_partition_mode}_to_{target_partition_mode}" + with _cp_layout_nvtx_range(f"cp_layout/thd/swap/{conversion_name}"): + if seq_dim != 0: + x = x.movedim(seq_dim, 0) + x = x.contiguous() + + route = thd_cp_partition_route + if route is None or route.device != x.device: + route = build_thd_cp_partition_route( + cu_seqlens, + cp_size, + cp_rank, + source_partition_mode, + target_partition_mode, + device=x.device, + ) + ( + local_source_length, + local_target_length, + send_rows, + recv_rows, + input_split_sizes, + output_split_sizes, + ) = decode_thd_cp_partition_route(route, cp_size, cp_rank) + + if x.size(0) != local_source_length: + raise ValueError( + f"Local THD tensor length ({x.size(0)}) does not match {source_partition_mode} " + f"rank-{cp_rank} partition length ({local_source_length})." + ) + + with _cp_layout_nvtx_range(f"cp_layout/thd/pack/{conversion_name}"): + send_buf = _pack_thd_cp_route_send_buffer(x, local_source_length, send_rows) + if not send_buf.is_contiguous(): + send_buf = send_buf.contiguous() + + with _cp_layout_nvtx_range(f"cp_layout/thd/all_to_all/{conversion_name}"): + recv_buf = all_to_all(cp_group, send_buf, output_split_sizes, input_split_sizes) + + with _cp_layout_nvtx_range(f"cp_layout/thd/scatter/{conversion_name}"): + out_shape = (local_target_length,) + tuple(x.shape[1:]) + out = _scatter_thd_cp_route_recv_buffer(recv_buf, recv_rows, out_shape) + + if seq_dim != 0: + out = out.movedim(0, seq_dim) + return out.contiguous() + + +def _zigzag_contiguous_chunk_swap( + x: torch.Tensor, + cp_group: Optional[torch.distributed.ProcessGroup], + seq_dim: int, + to_contiguous: bool, +) -> torch.Tensor: + """Single-all-to-all chunk permutation between zigzag and contiguous layouts. + + Each rank holds exactly two chunks along ``seq_dim``. The mapping from + local (rank, slot) to (rank, slot) in the target layout is deterministic + and depends only on ``cp_size`` and ``cp_rank``, so we pack send data in + destination-rank order and use one ``all_to_all_single`` with unequal + splits to route each chunk to its target rank. + """ + cp_size = cp_group.size() if cp_group is not None else 1 + if cp_size == 1: + return x + cp_rank = cp_group.rank() + from megatron.core.tensor_parallel.mappings import all_to_all + + # Work with seq_dim at position 0. + if seq_dim != 0: + x = x.movedim(seq_dim, 0) + x = x.contiguous() + + seq_len_local = x.size(0) + assert seq_len_local % 2 == 0, ( + f"zigzag/contiguous chunk swap requires an even local sequence length, " + f"got {seq_len_local}." + ) + chunk_len = seq_len_local // 2 + + def _rank_to_chunks(rank: int, in_zigzag: bool) -> Tuple[int, int]: + """Global chunk indices at (slot 0, slot 1) for this rank.""" + if in_zigzag: + return (rank, 2 * cp_size - rank - 1) + return (2 * rank, 2 * rank + 1) + + def _chunk_to_dest(chunk_idx: int, target_zigzag: bool) -> Tuple[int, int]: + """Destination (rank, slot) for a given global chunk index in the target layout.""" + if target_zigzag: + if chunk_idx < cp_size: + return chunk_idx, 0 + return 2 * cp_size - chunk_idx - 1, 1 + return chunk_idx // 2, chunk_idx % 2 + + # TODO(yuzhongw): cache this small SBHD permutation plan by + # (cp_size, cp_rank, source layout, target layout, device) instead of + # rebuilding Python lists on every conversion. + source_in_zigzag = to_contiguous + target_in_zigzag = not to_contiguous + source_partition_mode = "zigzag" if source_in_zigzag else "contiguous" + target_partition_mode = "zigzag" if target_in_zigzag else "contiguous" + conversion_name = f"{source_partition_mode}_to_{target_partition_mode}" + + local_chunk_indices = _rank_to_chunks(cp_rank, source_in_zigzag) + local_dests = [_chunk_to_dest(c, target_in_zigzag) for c in local_chunk_indices] + + # Pack the send buffer so chunks are ordered by (dst_rank, dst_slot). + local_slot_order = sorted(range(2), key=lambda s: local_dests[s]) + local_chunks = [x[:chunk_len], x[chunk_len:]] + send_buf = torch.cat([local_chunks[s] for s in local_slot_order], dim=0).contiguous() + + input_split_chunks = [0] * cp_size + for dst_rank, _ in local_dests: + input_split_chunks[dst_rank] += 1 + + # Mirror every source rank's packing logic so we know which received chunk + # belongs in which local target slot. + output_split_chunks = [0] * cp_size + recv_dst_slots_per_source: List[List[int]] = [[] for _ in range(cp_size)] + for src in range(cp_size): + src_chunks = _rank_to_chunks(src, source_in_zigzag) + src_dests = [_chunk_to_dest(c, target_in_zigzag) for c in src_chunks] + src_slot_order = sorted(range(2), key=lambda s: src_dests[s]) + for s in src_slot_order: + dst_rank, dst_slot = src_dests[s] + if dst_rank == cp_rank: + output_split_chunks[src] += 1 + recv_dst_slots_per_source[src].append(dst_slot) + + input_split_sizes = [n * chunk_len for n in input_split_chunks] + output_split_sizes = [n * chunk_len for n in output_split_chunks] + + with _cp_layout_nvtx_range(f"cp_layout/sbhd/all_to_all/{conversion_name}"): + recv_buf = all_to_all(cp_group, send_buf, output_split_sizes, input_split_sizes) + + # Reassemble local chunks in target-layout slot order. + target_slots: List[Optional[torch.Tensor]] = [None, None] + offset = 0 + for src in range(cp_size): + for dst_slot in recv_dst_slots_per_source[src]: + target_slots[dst_slot] = recv_buf[offset : offset + chunk_len] + offset += chunk_len + assert all(t is not None for t in target_slots), "Incomplete chunk reassembly in CP swap" + + out = torch.cat(target_slots, dim=0) + if seq_dim != 0: + out = out.movedim(0, seq_dim) + return out.contiguous() diff --git a/megatron/core/context_parallel_layout/metadata.py b/megatron/core/context_parallel_layout/metadata.py new file mode 100644 index 00000000000..e6a5d324b91 --- /dev/null +++ b/megatron/core/context_parallel_layout/metadata.py @@ -0,0 +1,38 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Packed-sequence metadata helpers for CP partition-mode tracking.""" + +from typing import Any, Optional + +import torch + +from megatron.core.context_parallel_layout import CpPartitionMode + + +def get_packed_seq_params_cp_partition_cu_seqlens( + packed_seq_params: Optional[Any], +) -> Optional[torch.Tensor]: + """Return THD cumulative sequence lengths used for CP layout conversion. + + SBHD callers may still pass synthetic ``PackedSeqParams`` for CP layout + annotation. Only THD metadata carries global packed-token boundaries. + """ + if packed_seq_params is None or getattr(packed_seq_params, "qkv_format", None) != "thd": + return None + return ( + packed_seq_params.cu_seqlens_q_padded + if packed_seq_params.cu_seqlens_q_padded is not None + else packed_seq_params.cu_seqlens_q + ) + + +def replace_packed_seq_params_cp_partition_mode( + packed_seq_params: Optional[Any], cp_partition_mode: Optional[CpPartitionMode] +) -> Optional[Any]: + """Annotate packed-sequence metadata with the current CP partition mode.""" + if packed_seq_params is None: + return packed_seq_params + if getattr(packed_seq_params, "cp_partition_mode", None) == cp_partition_mode: + return packed_seq_params + packed_seq_params.cp_partition_mode = cp_partition_mode + return packed_seq_params diff --git a/megatron/core/context_parallel_layout/policy.py b/megatron/core/context_parallel_layout/policy.py new file mode 100644 index 00000000000..858db99f66d --- /dev/null +++ b/megatron/core/context_parallel_layout/policy.py @@ -0,0 +1,124 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Layer and stage CP partition-mode policy helpers.""" + +from typing import Any, Optional + +import torch + +from megatron.core.context_parallel_layout import CpPartitionMode + + +def get_context_parallel_layout_chunk_indices( + cp_size: int, cp_rank: int, cp_partition_mode: str +) -> torch.Tensor: + """Return the two global chunk indices owned by this CP rank in ``cp_partition_mode``.""" + if cp_size < 1: + raise ValueError(f"cp_size must be >= 1, got {cp_size}.") + if not 0 <= cp_rank < cp_size: + raise ValueError(f"cp_rank must be in [0, {cp_size}), got {cp_rank}.") + + if cp_partition_mode == "zigzag": + return torch.tensor([cp_rank, 2 * cp_size - cp_rank - 1], dtype=torch.long) + if cp_partition_mode == "contiguous": + return torch.tensor([2 * cp_rank, 2 * cp_rank + 1], dtype=torch.long) + raise ValueError( + f"Unsupported context-parallel partition mode {cp_partition_mode!r} for " + f"cp_size={cp_size}, cp_rank={cp_rank}." + ) + + +################################################################################ +# Layer-to-CP-partition-mode mapping +################################################################################ +# ``None`` is a meaningful result here: it means the module is token-layout +# agnostic and preserves whichever CP partition mode it receives. It must not +# be used as the fallback for an unrecognized module type; unknown types should +# fail loudly so new layer implementations add an explicit partition-mode policy. + + +def _validate_cp_partition_mode_preference( + mode: Optional[str], *, module_name: str +) -> Optional[CpPartitionMode]: + if mode is None or mode in ("zigzag", "contiguous"): + return mode + raise ValueError(f"Invalid CP partition mode preference {mode!r} declared by {module_name}.") + + +def get_preferred_cp_partition_mode_for_layer( + layer: Any, config: Any, *, cp_comm_type: Optional[str] = None +) -> Optional[CpPartitionMode]: + """Return a layer/module's CP partition mode preference for auto layout rollout. + + Modules should declare their auto-mode layout preference via + ``get_preferred_cp_partition_mode()``. This is a preference rather than a + hard requirement: some modules can consume more than one layout but still + prefer one for performance or rollout consistency. Wrapper modules may + delegate to ``inner_layer`` or ``self_attention``. + """ + if cp_comm_type is None: + cp_comm_type = getattr(config, "cp_comm_type", None) + + if layer is None: + raise ValueError("Cannot determine CP partition mode for None.") + + module_name = layer.__class__.__name__ + get_preferred_mode = getattr(layer, "get_preferred_cp_partition_mode", None) + if callable(get_preferred_mode): + return _validate_cp_partition_mode_preference( + get_preferred_mode(), module_name=module_name + ) + + if hasattr(layer, "inner_layer"): + return get_preferred_cp_partition_mode_for_layer( + layer.inner_layer, getattr(layer, "config", config), cp_comm_type=cp_comm_type + ) + if hasattr(layer, "self_attention"): + return get_preferred_cp_partition_mode_for_layer( + layer.self_attention, getattr(layer, "config", config), cp_comm_type=cp_comm_type + ) + + raise ValueError(f"Cannot determine CP partition mode for layer/module type {module_name!r}.") + + +def get_stage_entry_partition_mode( + packed_seq_params: Optional[Any], + expected_stage_entry_partition_mode: Optional[CpPartitionMode], + *, + owner_name: str, + cp_group: Optional[Any] = None, +) -> Optional[CpPartitionMode]: + """Return and validate the CP partition mode at a stage input boundary.""" + expected_stage_entry_partition_mode = _validate_cp_partition_mode_preference( + expected_stage_entry_partition_mode, module_name=owner_name + ) + packed_partition_mode = _validate_cp_partition_mode_preference( + getattr(packed_seq_params, "cp_partition_mode", None), module_name=owner_name + ) + + stage_entry_partition_mode = ( + packed_partition_mode + if packed_partition_mode is not None + else expected_stage_entry_partition_mode + ) + if expected_stage_entry_partition_mode is not None and stage_entry_partition_mode is not None: + assert stage_entry_partition_mode == expected_stage_entry_partition_mode, ( + f"{owner_name} expected CP stage entry partition mode " + f"{expected_stage_entry_partition_mode!r}, but packed_seq_params carries " + f"{stage_entry_partition_mode!r}." + ) + + effective_cp_group = ( + cp_group if cp_group is not None else getattr(packed_seq_params, "cp_group", None) + ) + if ( + effective_cp_group is not None + and effective_cp_group.size() > 1 + and stage_entry_partition_mode is None + ): + raise ValueError( + f"{owner_name} requires a CP stage entry partition mode when context " + "parallelism is active." + ) + + return stage_entry_partition_mode diff --git a/megatron/core/context_parallel_layout/routes.py b/megatron/core/context_parallel_layout/routes.py new file mode 100644 index 00000000000..10f662088d4 --- /dev/null +++ b/megatron/core/context_parallel_layout/routes.py @@ -0,0 +1,514 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""THD context-parallel partition indices and route tensor helpers.""" + +import warnings +from contextlib import contextmanager +from typing import Any, List, Optional, Tuple + +import torch + +from megatron.core.context_parallel_layout.metadata import ( + get_packed_seq_params_cp_partition_cu_seqlens, +) +from megatron.core.context_parallel_layout import CpPartitionMode + + +_THD_CP_ROUTE_HEADER_SIZE = 6 +_THD_CP_ROUTE_ATTRS = { + ("zigzag", "contiguous"): "cp_partition_route_zigzag_to_contiguous", + ("contiguous", "zigzag"): "cp_partition_route_contiguous_to_zigzag", +} + + +@contextmanager +def _cp_layout_nvtx_range(message: str): + active = torch.cuda.is_available() + if active: + torch.cuda.nvtx.range_push(message) + try: + yield + finally: + if active: + torch.cuda.nvtx.range_pop() + + +def get_thd_context_parallel_rank_indices( + cu_seqlens: torch.Tensor, cp_size: int, cp_rank: int, cp_partition_mode: str +) -> torch.Tensor: + """Return global THD token indices owned by one CP rank in a layout. + + Args: + cu_seqlens: Global packed-sequence cumulative lengths before CP partitioning. + cp_size: Context-parallel group size. + cp_rank: Context-parallel rank. + cp_partition_mode: Either ``"zigzag"`` or ``"contiguous"``. + + The returned indices are ordered exactly as the rank-local THD tensor is stored. + ``"zigzag"`` follows Megatron's per-sequence load-balanced chunk order; ``"contiguous"`` + partitions the flattened packed THD buffer into rank-contiguous spans. + """ + if cp_size < 1: + raise ValueError(f"cp_size must be >= 1, got {cp_size}.") + if not 0 <= cp_rank < cp_size: + raise ValueError(f"cp_rank must be in [0, {cp_size}), got {cp_rank}.") + if cu_seqlens.dim() != 1: + raise ValueError(f"cu_seqlens must be 1-D, got shape {tuple(cu_seqlens.shape)}.") + + cu = cu_seqlens.to(dtype=torch.long) + if cu.numel() == 0 or cu[0].item() != 0: + raise ValueError(f"cu_seqlens must start at 0, got {cu_seqlens}.") + + if torch.any(torch.diff(cu) < 0): + raise ValueError(f"cu_seqlens must be nondecreasing, got {cu_seqlens}.") + + nonduplicate_boundaries = torch.ones(cu.numel(), device=cu.device, dtype=torch.bool) + nonduplicate_boundaries[1:] = cu[1:] != cu[:-1] + cu = cu[nonduplicate_boundaries] + + total_tokens = int(cu[-1].item()) + if cp_partition_mode == "contiguous": + if total_tokens % cp_size != 0: + raise ValueError( + f"Contiguous CP partitioning requires total_tokens={total_tokens} " + f"to be divisible by cp_size={cp_size}." + ) + part_len = total_tokens // cp_size + rank_start = cp_rank * part_len + return torch.arange(rank_start, rank_start + part_len, device=cu.device, dtype=torch.long) + if cp_partition_mode != "zigzag": + raise ValueError( + f"Unsupported context-parallel partition mode {cp_partition_mode!r} " + f"for THD rank indices with cp_size={cp_size}, cp_rank={cp_rank}, " + f"cu_seqlens_shape={tuple(cu_seqlens.shape)}." + ) + + positions = torch.arange(total_tokens, device=cu.device, dtype=torch.long) + if total_tokens == 0: + return positions + + seq_lens = torch.diff(cu) + + chunk_divisor = 2 * cp_size + if torch.any(seq_lens % chunk_divisor != 0): + raise ValueError( + "All packed sequence lengths must be divisible by " + f"2 * cp_size ({chunk_divisor}) for zigzag CP layout conversion, " + f"got {seq_lens}." + ) + + seq_idx = torch.bucketize(positions, cu[1:], right=True) + global_starts = cu[:-1] + pos_in_seq = positions - global_starts[seq_idx] + chunk_lens = (seq_lens // chunk_divisor)[seq_idx] + chunk = pos_in_seq // chunk_lens + offset = pos_in_seq - chunk * chunk_lens + + owner = torch.where(chunk < cp_size, chunk, 2 * cp_size - chunk - 1) + local_slot = torch.where(chunk < cp_size, torch.zeros_like(chunk), torch.ones_like(chunk)) + + local_starts = (global_starts // cp_size)[seq_idx] + local_pos = local_starts + local_slot * chunk_lens + offset + + rank_mask = owner == cp_rank + rank_positions = positions[rank_mask] + rank_local_pos = local_pos[rank_mask] + return rank_positions[torch.argsort(rank_local_pos)] + + +_ThdLayoutSegment = Tuple[int, int, int] + + +def _compact_thd_cu_seqlens_to_list(cu_seqlens: torch.Tensor) -> List[int]: + if cu_seqlens.dim() != 1: + raise ValueError(f"cu_seqlens must be 1-D, got shape {tuple(cu_seqlens.shape)}.") + + cu = cu_seqlens.detach().to(device="cpu", dtype=torch.long).tolist() + if not cu or cu[0] != 0: + raise ValueError(f"cu_seqlens must start at 0, got {cu_seqlens}.") + + compact_cu: List[int] = [cu[0]] + prev = cu[0] + for value in cu[1:]: + if value < prev: + raise ValueError(f"cu_seqlens must be nondecreasing, got {cu_seqlens}.") + if value != prev: + compact_cu.append(value) + prev = value + return compact_cu + + +def _validate_thd_route_partitioning(cu: List[int], cp_size: int) -> None: + total_tokens = cu[-1] + if total_tokens % cp_size != 0: + raise ValueError( + f"Contiguous CP partitioning requires total_tokens={total_tokens} " + f"to be divisible by cp_size={cp_size}." + ) + + chunk_divisor = 2 * cp_size + bad_seq_lens = [ + seq_end - seq_start + for seq_start, seq_end in zip(cu[:-1], cu[1:]) + if (seq_end - seq_start) % chunk_divisor != 0 + ] + if bad_seq_lens: + raise ValueError( + "All packed sequence lengths must be divisible by " + f"2 * cp_size ({chunk_divisor}) for zigzag CP layout conversion, " + f"got {bad_seq_lens}." + ) + + +def _build_thd_layout_segments( + cu: List[int], cp_size: int, cp_rank: int, cp_partition_mode: CpPartitionMode +) -> Tuple[List[_ThdLayoutSegment], int]: + total_tokens = cu[-1] + if cp_partition_mode == "contiguous": + part_len = total_tokens // cp_size + if part_len == 0: + return [], 0 + return [(cp_rank * part_len, part_len, 0)], part_len + + if cp_partition_mode != "zigzag": + raise ValueError( + f"Unsupported context-parallel partition mode {cp_partition_mode!r} " + f"for THD layout segments with cp_size={cp_size}, rank={cp_rank}." + ) + + segments: List[_ThdLayoutSegment] = [] + local_start = 0 + for seq_start, seq_end in zip(cu[:-1], cu[1:]): + seq_len = seq_end - seq_start + chunk_len = seq_len // (2 * cp_size) + first_chunk = cp_rank + second_chunk = 2 * cp_size - cp_rank - 1 + segments.append((seq_start + first_chunk * chunk_len, chunk_len, local_start)) + segments.append((seq_start + second_chunk * chunk_len, chunk_len, local_start + chunk_len)) + local_start += 2 * chunk_len + + return segments, local_start + + +def _intersect_thd_layout_segments( + source_segments: List[_ThdLayoutSegment], target_segments: List[_ThdLayoutSegment] +) -> List[Tuple[int, int, int]]: + intersections: List[Tuple[int, int, int]] = [] + source_index = 0 + target_index = 0 + while source_index < len(source_segments) and target_index < len(target_segments): + source_global_start, source_len, source_local_start = source_segments[source_index] + target_global_start, target_len, target_local_start = target_segments[target_index] + source_global_end = source_global_start + source_len + target_global_end = target_global_start + target_len + + overlap_start = max(source_global_start, target_global_start) + overlap_end = min(source_global_end, target_global_end) + if overlap_start < overlap_end: + intersections.append( + ( + source_local_start + overlap_start - source_global_start, + target_local_start + overlap_start - target_global_start, + overlap_end - overlap_start, + ) + ) + + if source_global_end <= target_global_end: + source_index += 1 + else: + target_index += 1 + + return intersections + + +def _append_range(rows: List[int], start: int, length: int) -> None: + rows.extend(range(start, start + length)) + + +def _row_list_is_identity(rows: List[int]) -> bool: + return all(row == index for index, row in enumerate(rows)) + + +def _thd_cp_partition_route_attr_name( + source_partition_mode: CpPartitionMode, target_partition_mode: CpPartitionMode +) -> str: + try: + return _THD_CP_ROUTE_ATTRS[(source_partition_mode, target_partition_mode)] + except KeyError as exc: + raise ValueError( + f"Unsupported CP partition mode conversion " + f"{source_partition_mode!r} -> {target_partition_mode!r} for THD route." + ) from exc + + +def _encode_thd_cp_partition_route( + *, + cp_size: int, + cp_rank: int, + local_source_length: int, + local_target_length: int, + send_rows_list: List[int], + recv_rows_list: List[int], + input_split_sizes: List[int], + output_split_sizes: List[int], + device: torch.device, +) -> torch.Tensor: + send_rows_payload = [] if _row_list_is_identity(send_rows_list) else send_rows_list + recv_rows_payload = [] if _row_list_is_identity(recv_rows_list) else recv_rows_list + payload = ( + [ + cp_size, + cp_rank, + local_source_length, + local_target_length, + len(send_rows_payload), + len(recv_rows_payload), + ] + + input_split_sizes + + output_split_sizes + + send_rows_payload + + recv_rows_payload + ) + return torch.tensor(payload, device=device, dtype=torch.long) + + +def _split_sizes_from_route_tensor( + route_tensor: torch.Tensor, start: int, end: int +) -> List[int]: + return route_tensor[start:end].detach().to(device="cpu", dtype=torch.long).tolist() + + +def decode_thd_cp_partition_route( + route_tensor: torch.Tensor, cp_size: int, cp_rank: int +) -> Tuple[int, int, Optional[torch.Tensor], Optional[torch.Tensor], List[int], List[int]]: + """Decode a THD CP route tensor into local conversion metadata. + + The tensor layout is: + ``[cp_size, cp_rank, local_source_len, local_target_len, send_rows_len, + recv_rows_len, input_splits..., output_splits..., send_rows..., recv_rows...]``. + Empty send/recv row payloads denote identity row order. + """ + if route_tensor is None: + raise ValueError("THD CP partition route tensor must not be None.") + if route_tensor.dim() != 1: + raise ValueError( + f"THD CP partition route tensor must be 1-D, got shape {tuple(route_tensor.shape)}." + ) + if route_tensor.numel() < _THD_CP_ROUTE_HEADER_SIZE: + raise ValueError( + f"THD CP partition route tensor is too short: {route_tensor.numel()} values." + ) + + ( + route_cp_size, + route_cp_rank, + local_source_length, + local_target_length, + send_rows_len, + recv_rows_len, + ) = route_tensor[:_THD_CP_ROUTE_HEADER_SIZE].detach().cpu().tolist() + route_cp_size = int(route_cp_size) + route_cp_rank = int(route_cp_rank) + local_source_length = int(local_source_length) + local_target_length = int(local_target_length) + send_rows_len = int(send_rows_len) + recv_rows_len = int(recv_rows_len) + if route_cp_size != cp_size or route_cp_rank != cp_rank: + raise ValueError( + "THD CP partition route tensor does not match the requested CP rank/size: " + f"route cp_size={route_cp_size}, cp_rank={route_cp_rank}; " + f"requested cp_size={cp_size}, cp_rank={cp_rank}." + ) + split_start = _THD_CP_ROUTE_HEADER_SIZE + input_split_start = split_start + output_split_start = input_split_start + cp_size + send_rows_start = output_split_start + cp_size + recv_rows_start = send_rows_start + send_rows_len + expected_numel = recv_rows_start + recv_rows_len + if route_tensor.numel() != expected_numel: + raise ValueError( + "THD CP partition route tensor has inconsistent length: " + f"got {route_tensor.numel()}, expected {expected_numel}." + ) + + input_split_sizes = _split_sizes_from_route_tensor( + route_tensor, input_split_start, output_split_start + ) + output_split_sizes = _split_sizes_from_route_tensor( + route_tensor, output_split_start, send_rows_start + ) + send_rows = None if send_rows_len == 0 else route_tensor[send_rows_start:recv_rows_start] + recv_rows = None if recv_rows_len == 0 else route_tensor[recv_rows_start:expected_numel] + return ( + local_source_length, + local_target_length, + send_rows, + recv_rows, + input_split_sizes, + output_split_sizes, + ) + + +def build_thd_cp_partition_route( + cu_seqlens: torch.Tensor, + cp_size: int, + cp_rank: int, + source_partition_mode: CpPartitionMode, + target_partition_mode: CpPartitionMode, + *, + device: Optional[torch.device] = None, +) -> torch.Tensor: + """Precompute one THD CP layout conversion route as a tensor. + + The route depends only on packed sequence metadata, CP rank/size, and the + source/target partition modes. It can be reused for every tensor that has + the same THD sequence axis in the same microbatch. + """ + if source_partition_mode not in ("zigzag", "contiguous") or target_partition_mode not in ( + "zigzag", + "contiguous", + ): + raise ValueError( + f"Unsupported CP partition mode conversion " + f"{source_partition_mode!r} -> {target_partition_mode!r} for THD route: " + f"cp_size={cp_size}, cp_rank={cp_rank}, cu_seqlens_shape={tuple(cu_seqlens.shape)}." + ) + if source_partition_mode == target_partition_mode: + raise ValueError("A THD CP partition route is only needed when partition modes differ.") + _thd_cp_partition_route_attr_name(source_partition_mode, target_partition_mode) + if device is None: + device = cu_seqlens.device + + with _cp_layout_nvtx_range( + f"cp_layout/thd/route/{source_partition_mode}_to_{target_partition_mode}" + ): + cu = _compact_thd_cu_seqlens_to_list(cu_seqlens) + _validate_thd_route_partitioning(cu, cp_size) + + source_segments_by_rank: List[List[_ThdLayoutSegment]] = [] + source_lengths: List[int] = [] + target_segments_by_rank: List[List[_ThdLayoutSegment]] = [] + target_lengths: List[int] = [] + for rank in range(cp_size): + source_segments, source_length = _build_thd_layout_segments( + cu, cp_size, rank, source_partition_mode + ) + target_segments, target_length = _build_thd_layout_segments( + cu, cp_size, rank, target_partition_mode + ) + source_segments_by_rank.append(source_segments) + source_lengths.append(source_length) + target_segments_by_rank.append(target_segments) + target_lengths.append(target_length) + + local_source_segments = source_segments_by_rank[cp_rank] + local_target_segments = target_segments_by_rank[cp_rank] + + send_rows_list: List[int] = [] + input_split_sizes: List[int] = [] + for dst_rank in range(cp_size): + intersections = _intersect_thd_layout_segments( + local_source_segments, target_segments_by_rank[dst_rank] + ) + intersections.sort(key=lambda item: item[1]) + input_split_size = 0 + for source_row, _, length in intersections: + _append_range(send_rows_list, source_row, length) + input_split_size += length + input_split_sizes.append(input_split_size) + + recv_rows_list: List[int] = [] + output_split_sizes: List[int] = [] + for src_rank in range(cp_size): + intersections = _intersect_thd_layout_segments( + source_segments_by_rank[src_rank], local_target_segments + ) + intersections.sort(key=lambda item: item[1]) + output_split_size = 0 + for _, target_row, length in intersections: + _append_range(recv_rows_list, target_row, length) + output_split_size += length + output_split_sizes.append(output_split_size) + + assert len(send_rows_list) == source_lengths[cp_rank] + assert len(recv_rows_list) == target_lengths[cp_rank] + return _encode_thd_cp_partition_route( + cp_size=cp_size, + cp_rank=cp_rank, + local_source_length=source_lengths[cp_rank], + local_target_length=target_lengths[cp_rank], + send_rows_list=send_rows_list, + recv_rows_list=recv_rows_list, + input_split_sizes=input_split_sizes, + output_split_sizes=output_split_sizes, + device=device, + ) + + +def get_thd_cp_partition_route( + packed_seq_params: Optional[Any], + source_partition_mode: CpPartitionMode, + target_partition_mode: CpPartitionMode, +) -> Optional[torch.Tensor]: + """Return the precomputed THD CP partition route tensor for one direction.""" + if source_partition_mode == target_partition_mode: + return None + if packed_seq_params is None or getattr(packed_seq_params, "qkv_format", None) != "thd": + return None + + attr_name = _thd_cp_partition_route_attr_name(source_partition_mode, target_partition_mode) + route = getattr(packed_seq_params, attr_name, None) + if route is not None: + return route + + warnings.warn( + "THD PackedSeqParams is missing precomputed context-parallel layout routes. " + "This lookup will attempt to build them from packed_seq_params.cp_group as " + "a compatibility fallback. Callers should prebuild THD CP routes when " + "constructing the batch; a future release will require the routes to be " + "present before layout conversion.", + FutureWarning, + stacklevel=2, + ) + prebuild_thd_cp_partition_routes(packed_seq_params) + return getattr(packed_seq_params, attr_name, None) + + +def prebuild_thd_cp_partition_routes( + packed_seq_params: Optional[Any], + cp_group: Optional[torch.distributed.ProcessGroup] = None, + *, + device: Optional[torch.device] = None, +) -> None: + """Best-effort prebuild of THD CP layout route tensors for a packed microbatch.""" + if packed_seq_params is None or getattr(packed_seq_params, "qkv_format", None) != "thd": + return + if cp_group is None: + cp_group = getattr(packed_seq_params, "cp_group", None) + if cp_group is None or cp_group.size() <= 1: + return + cp_size = cp_group.size() + cp_rank = cp_group.rank() + cu_seqlens = get_packed_seq_params_cp_partition_cu_seqlens(packed_seq_params) + if cu_seqlens is None: + return + if device is None: + device = cu_seqlens.device + + for source_partition_mode, target_partition_mode in _THD_CP_ROUTE_ATTRS: + attr_name = _thd_cp_partition_route_attr_name(source_partition_mode, target_partition_mode) + try: + route = build_thd_cp_partition_route( + cu_seqlens, + cp_size, + cp_rank, + source_partition_mode, + target_partition_mode, + device=device, + ) + except ValueError: + # Some batches/layouts may never need the opposite route. Preserve + # lazy block-time validation for the path that actually uses it. + setattr(packed_seq_params, attr_name, None) + continue + setattr(packed_seq_params, attr_name, route) diff --git a/megatron/core/datasets/data_schedule.py b/megatron/core/datasets/data_schedule.py index b6790c01d8a..a7ad598162a 100644 --- a/megatron/core/datasets/data_schedule.py +++ b/megatron/core/datasets/data_schedule.py @@ -573,6 +573,7 @@ def get_batch_on_this_rank_for_sequence_packing( dynamic_cp: bool = False, pg_collection: Optional[ProcessGroupCollection] = None, config=None, + cp_partition_mode: Optional[str] = None, ): """ Get a batch of data for sequence packing. @@ -580,8 +581,10 @@ def get_batch_on_this_rank_for_sequence_packing( data_iterator (Iterator): The data iterator to get the batch from. mtp_on_this_rank (bool): Whether to use multi-token prediction. vp_stage (Optional[int]): The stage of the pipeline. - config: Model config used for CP partitioning and optional THD packed-sequence padding. - When None, CP partitioning defaults to zigzag and no padding is applied. + config: Model config used for optional THD packed-sequence padding. + When None, padding is disabled. + cp_partition_mode: CP partition mode requested by the current model chunk input. + Required when this batch is partitioned across multiple CP ranks. Returns: tuple of (tokens, labels, loss_mask, attention_mask, position_ids, packed_seq_params, padding_mask) @@ -637,7 +640,12 @@ def get_batch_on_this_rank_for_sequence_packing( group_size=local_cp_size_val ) - cp_partition_mode = getattr(config, "cp_partition_mode", "zigzag") + if cp_partition_mode is None: + if cp_group.size() > 1: + raise ValueError( + "cp_partition_mode must be provided when sequence-packing batches are " + "partitioned across multiple context-parallel ranks." + ) tail_padding_policy = resolve_thd_tail_padding_policy(config) contiguous_cp_local_target_len = None non_dummy_global_target_len = None diff --git a/megatron/core/datasets/data_schedule_utils.py b/megatron/core/datasets/data_schedule_utils.py index 120513c1819..2e3538be2b5 100644 --- a/megatron/core/datasets/data_schedule_utils.py +++ b/megatron/core/datasets/data_schedule_utils.py @@ -16,7 +16,7 @@ def get_cp_slice_for_thd( batch, cp_group, keys: Optional[Sequence[str]] = None, - cp_partition_mode: Literal["zigzag", "contiguous"] = "zigzag", + cp_partition_mode: Optional[Literal["zigzag", "contiguous"]] = None, partition_total_tokens: Optional[int] = None, ): """Partition sequence data for context parallelism in THD format. @@ -36,6 +36,8 @@ def get_cp_slice_for_thd( cp_size = cp_group.size() if cp_size <= 1 and partition_total_tokens is None: return + if cp_partition_mode is None: + raise ValueError("cp_partition_mode must be provided for THD context parallel slicing.") cp_rank = cp_group.rank() # Partition with padded cumulative lengths so CP slices match the THD # sequence boundaries consumed by attention kernels. diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 8d797e816db..f9f7500cc42 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -26,7 +26,7 @@ maybe_fake_quantize_int4_weight_tensors, ) from megatron.core.model_parallel_config import ModelParallelConfig -from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.packed_seq_params import PackedSeqParams, THD_CP_PARTITION_ROUTE_TENSOR_FIELDS from megatron.core.parallel_state import ( get_amax_reduction_group, get_context_parallel_group, @@ -1758,6 +1758,8 @@ def __init__( self.kept_packed_seq_params.discard("seq_idx") self.kept_packed_seq_params.discard("tokens_per_sample") self.kept_packed_seq_params.discard("cp_partition_mode") + for field_name in THD_CP_PARTITION_ROUTE_TENSOR_FIELDS: + self.kept_packed_seq_params.discard(field_name) if config.qk_clip or config.log_max_attention_logit: # qk-clip is only supported in TE 2.9.0 and later diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index 3db7c9987b1..226f9d7a6f2 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -2,6 +2,7 @@ from typing import List, Optional +from megatron.core.context_parallel_layout import CpPartitionMode from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add from megatron.core.models.backends import BackendSpecProvider from megatron.core.ssm.gated_delta_net import GatedDeltaNet, GatedDeltaNetSubmodules @@ -453,6 +454,77 @@ def _validate_dsa_index_share_pipeline_split(config: TransformerConfig, local_la ) +def get_experimental_attention_variant_layer_cp_partition_mode_pattern( + config: TransformerConfig, +) -> List[Optional[CpPartitionMode]]: + """Return per-layer CP partition-mode preferences for experimental GPT specs.""" + if config.experimental_attention_variant is None: + return ["zigzag"] * config.num_layers + + experimental_attention_pattern = [0] * config.num_layers + if is_linear_attention_variant(config.experimental_attention_variant): + experimental_attention_pattern = get_linear_attention_pattern(config=config) + else: + experimental_attention_pattern = [1] * config.num_layers + + experimental_layout = _get_experimental_attention_variant_cp_partition_mode(config) + return [ + experimental_layout if uses_experimental_attention else "zigzag" + for uses_experimental_attention in experimental_attention_pattern + ] + + +def get_experimental_attention_variant_stage_input_cp_partition_mode( + config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None +) -> Optional[CpPartitionMode]: + """Return the CP partition mode expected at one experimental GPT stage input.""" + layer_layouts = get_experimental_attention_variant_layer_cp_partition_mode_pattern(config) + + if config.pipeline_model_parallel_layout is not None: + stage_layer_offset = config.pipeline_model_parallel_layout.get_layer_offset( + layer_type=LayerType.decoder, vp_stage=vp_stage, pp_rank=pp_rank + ) + elif config.pipeline_model_parallel_size == 1 and pp_rank is None: + stage_layer_offset = 0 + else: + stage_layer_offset = get_transformer_layer_offset( + config, vp_stage=vp_stage, pp_rank=pp_rank + ) + + current_partition_mode = None + for preferred_partition_mode in layer_layouts[:stage_layer_offset]: + if preferred_partition_mode is not None: + current_partition_mode = preferred_partition_mode + + if current_partition_mode is None: + for preferred_partition_mode in layer_layouts[stage_layer_offset:]: + if preferred_partition_mode is not None: + current_partition_mode = preferred_partition_mode + break + + return current_partition_mode + + +def _get_experimental_attention_variant_cp_partition_mode( + config: TransformerConfig, +) -> CpPartitionMode: + """Return the CP partition mode preferred by the experimental attention variant.""" + if config.experimental_attention_variant == "gated_delta_net": + mode = getattr(config, "linear_cp_mode", "chunkwise") + if mode == "chunkwise": + return "contiguous" + if mode == "headwise": + return "zigzag" + raise ValueError(f"Unsupported GatedDeltaNet linear_cp_mode: {mode!r}.") + if config.experimental_attention_variant == "dsv4_hybrid": + return "contiguous" + if config.experimental_attention_variant == "dsa": + return "zigzag" + raise ValueError( + f"Invalid experimental attention variant: {config.experimental_attention_variant}" + ) + + def get_moe_layer_pattern(config: TransformerConfig) -> List[int]: """Parse config.moe_layer_freq to get per-layer MoE pattern (1=MoE, 0=dense). diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index a986463bcde..0133b74a4e5 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -1,5 +1,6 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import warnings from collections import OrderedDict from typing import Any, Callable, Dict, Literal, Optional @@ -8,6 +9,11 @@ from megatron.core import tensor_parallel from megatron.core.config_logger import has_config_logger_enabled, log_config_to_disk +from megatron.core.context_parallel_layout import ( + CpPartitionModeConverter, + get_stage_entry_partition_mode, + replace_packed_seq_params_cp_partition_mode, +) from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.extensions.transformer_engine import TELMHeadColumnParallelLinear from megatron.core.fp8_utils import is_mxfp8_output_proj_active @@ -111,6 +117,7 @@ def __init__( mtp_block_spec: Optional[ModuleSpec] = None, pg_collection: Optional[ProcessGroupCollection] = None, vp_stage: Optional[int] = None, + cp_stage_entry_partition_mode: Optional[str] = None, ) -> None: super().__init__(config=config, pg_collection=pg_collection) @@ -235,6 +242,7 @@ def __init__( post_process=self.post_process, pg_collection=self.pg_collection, vp_stage=vp_stage, + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, ) if self.mtp_process: @@ -325,6 +333,7 @@ def _preprocess( inference_context: BaseInferenceContext = None, packed_seq_params: PackedSeqParams = None, padding_mask: Optional[Tensor] = None, + input_partition_mode=None, ): """Preprocesses inputs for the transformer decoder. @@ -372,6 +381,9 @@ def _preprocess( rotary_pos_sin = None # this is used to store combined cos/sin embeddings, exclusively for flash infer rope rotary_pos_cos_sin = None + # Model-level rotary_pos_emb is only for regular attention. Regular + # attention uses the default zigzag CP RoPE layout; MLA/CSA/DSv4-style + # variants must ignore this external RoPE and build/apply RoPE internally. if self.position_embedding_type == 'rope' and not self.config.multi_latent_attention: use_flash_infer_fused_rope = ( @@ -579,6 +591,27 @@ def forward( inference_context = deprecate_inference_params(inference_context, inference_params) + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + if cp_group is not None and cp_group.size() > 1 and packed_seq_params is None: + warnings.warn( + "GPTModel received no PackedSeqParams while running under context " + "parallelism. Megatron-LM will temporarily assume SBHD tensors and " + "create layout metadata for this forward pass. In a future release, " + "callers must pass PackedSeqParams with qkv_format and cp_partition_mode " + "set explicitly.", + FutureWarning, + stacklevel=2, + ) + packed_seq_params = PackedSeqParams( + qkv_format="sbhd", cp_partition_mode=self.decoder.cp_stage_entry_partition_mode + ) + input_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + self.decoder.cp_stage_entry_partition_mode, + owner_name=type(self).__name__, + cp_group=cp_group, + ) + preproc_output = self._preprocess( input_ids=input_ids, position_ids=position_ids, @@ -586,6 +619,7 @@ def forward( inference_context=inference_context, packed_seq_params=packed_seq_params, padding_mask=padding_mask, + input_partition_mode=input_partition_mode, ) ( @@ -642,6 +676,7 @@ def forward( inference_params=inference_params, packed_seq_params=packed_seq_params, sequence_len_offset=sequence_len_offset, + input_partition_mode=input_partition_mode, runtime_gather_output=runtime_gather_output, extra_block_kwargs=extra_block_kwargs, inference_context=inference_context, @@ -667,6 +702,7 @@ def _postprocess( inference_params=None, packed_seq_params=None, sequence_len_offset=None, + input_partition_mode=None, runtime_gather_output=None, extra_block_kwargs=None, inference_context=None, @@ -697,7 +733,77 @@ def _postprocess( output_weight = None if self.share_embeddings_and_output_weights: output_weight = self.shared_embedding_or_output_weight() - if mtp_in_postprocess and not (in_inference_mode or is_spec_decode): + + postprocess_to_input_converter = None + mtp_forward_ran = mtp_in_postprocess and not (in_inference_mode or is_spec_decode) + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + cp_size = cp_group.size() if cp_group is not None else 1 + + needs_batch_layout = ( + cp_size > 1 + and self.config.cp_partition_mode == "auto" + and (self.post_process or mtp_forward_ran) + ) + if needs_batch_layout: + if input_partition_mode is None: + input_partition_mode = self.decoder.cp_stage_entry_partition_mode + block_output_partition_mode = getattr( + packed_seq_params, "cp_partition_mode", input_partition_mode + ) + postprocess_partition_mode = ( + block_output_partition_mode if mtp_forward_ran else input_partition_mode + ) + block_to_postprocess_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=block_output_partition_mode, + target_partition_mode=postprocess_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + hidden_states = block_to_postprocess_converter.convert( + hidden_states, seq_dim=0, sequence_parallel=self.config.sequence_parallel + ) + if mhc_multistream is not None: + mhc_multistream = block_to_postprocess_converter.convert( + mhc_multistream, seq_dim=0, sequence_parallel=self.config.sequence_parallel + ) + if input_partition_mode != postprocess_partition_mode: + input_to_postprocess_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=input_partition_mode, + target_partition_mode=postprocess_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + input_to_postprocess_converter.assert_no_dense_attention_inputs( + attention_mask=attention_mask, + hidden_states=hidden_states, + ) + input_ids = input_to_postprocess_converter.convert(input_ids, seq_dim=-1) + position_ids = input_to_postprocess_converter.convert(position_ids, seq_dim=-1) + labels = input_to_postprocess_converter.convert(labels, seq_dim=-1) + loss_mask = input_to_postprocess_converter.convert(loss_mask, seq_dim=-1) + padding_mask = input_to_postprocess_converter.convert( + padding_mask, seq_dim=-1, sequence_parallel=self.config.sequence_parallel + ) + # Model-level rotary_pos_emb belongs to regular attention, whose + # CP layout preference is zigzag. MTP side tensors are aligned + # for token/loss semantics, but RoPE is not treated as a batch + # side tensor to be converted here. + postprocess_to_input_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=postprocess_partition_mode, + target_partition_mode=input_partition_mode, + config=self.config, + ) + packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, postprocess_partition_mode + ) + + if mtp_forward_ran: hidden_states = self.mtp( input_ids=input_ids, position_ids=position_ids, @@ -821,6 +927,8 @@ def _postprocess( if labels is None: # [s b h] => [b s h] + if postprocess_to_input_converter is not None: + logits = postprocess_to_input_converter.convert(logits, seq_dim=0) return logits.transpose(0, 1).contiguous() output_layer_kwargs = dict( @@ -836,6 +944,11 @@ def _postprocess( logits, _ = self.output_layer(**output_layer_kwargs) loss = self.compute_language_model_loss(labels, logits) + if postprocess_to_input_converter is not None: + # The training loss function masks this returned per-token loss with the + # original batch loss_mask, so preserve the model input layout at the boundary. + loss = postprocess_to_input_converter.convert(loss, seq_dim=-1) + return loss def build_schedule_plan( diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index 297bfc4b054..d417d21cac1 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -12,6 +12,13 @@ import torch from torch import Tensor, nn +from megatron.core.context_parallel_layout import ( + CpPartitionMode, + CpPartitionModeConverter, + get_preferred_cp_partition_mode_for_layer, + get_stage_entry_partition_mode, + replace_packed_seq_params_cp_partition_mode, +) from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import replace_prefix_for_sharding from megatron.core.enums import Fp8Recipe @@ -21,7 +28,7 @@ from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.inference.utils import InferenceMode from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols as LayerSymbols -from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.recompute import checkpointed_forward from megatron.core.tensor_parallel.random import CheckpointManager @@ -160,9 +167,17 @@ def _reconstruct_packed_seq_params_from_kwargs(self, kwargs): if 'cu_seqlens_q' not in kwargs: return max_seqlen = self.config.max_seqlen_per_dp_cp_rank * self.config.context_parallel_size + cp_partition_mode = get_preferred_cp_partition_mode_for_layer(self, self.config) + if cp_partition_mode is None: + if self.config.context_parallel_size > 1 or self.config.dynamic_context_parallel: + raise ValueError( + "Cannot reconstruct THD PackedSeqParams for a layout-agnostic Hybrid layer " + "under context parallelism. The CP partition mode must be provided by the " + "model-level layout plan." + ) packed_seq_params = PackedSeqParams( qkv_format='thd', - cp_partition_mode=self.config.cp_partition_mode, + cp_partition_mode=cp_partition_mode, cu_seqlens_q=kwargs.pop('cu_seqlens_q'), cu_seqlens_kv=kwargs.pop('cu_seqlens_kv'), cu_seqlens_q_padded=kwargs.pop('cu_seqlens_q_padded'), @@ -412,6 +427,7 @@ def _call_inner_transformer_layer_without_local_bda( inference_context=inference_context, padding_mask=padding_mask, input_ids=input_ids, + packed_seq_params=packed_seq_params, ) if layer.mlp_norm_manager is not None: output_with_bias = layer._group_offload_output_with_bias( @@ -546,6 +562,8 @@ class HybridStack(MegatronModule): process groups to use. is_mtp_layer (bool, optional): whether this is an MTP layer. Defaults to False. mtp_layer_number (int, optional): enclosing MTP depth for logging nested MTP metrics. + cp_stage_entry_partition_mode (str, optional): CP partition mode expected at this + stage input. Required when context parallelism is enabled. """ def __init__( @@ -562,6 +580,7 @@ def __init__( pg_collection: ProcessGroupCollection = None, is_mtp_layer: bool = False, mtp_layer_number: Optional[int] = None, + cp_stage_entry_partition_mode: Optional[str] = None, name: str | None = None, ) -> None: """ @@ -574,6 +593,7 @@ def __init__( self.post_process = post_process self.is_mtp_layer = is_mtp_layer self.mtp_layer_number = mtp_layer_number + self.cp_stage_entry_partition_mode = cp_stage_entry_partition_mode assert pg_collection is not None, "pg_collection must be provided for HybridStack" @@ -747,6 +767,63 @@ def _set_mtp_layer_number_for_moe_metrics( if router is not None and getattr(router, "is_mtp_layer", False): router.mtp_layer_number = mtp_layer_number + def _convert_cp_partition_mode_for_layer( + self, + *, + local_index: int, + current_partition_mode: CpPartitionMode, + hidden_states: Tensor, + attention_mask: Optional[Tensor], + packed_seq_params: Optional[PackedSeqParams], + padding_mask: Optional[Tensor], + input_ids: Optional[Tensor], + preferred_partition_mode: Optional[CpPartitionMode], + ): + """Convert per-token tensors to the layout preferred by one local hybrid layer.""" + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + if cp_group is None or cp_group.size() <= 1: + return hidden_states, padding_mask, input_ids + if preferred_partition_mode is None or preferred_partition_mode == current_partition_mode: + return hidden_states, padding_mask, input_ids + + current_to_preferred_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=current_partition_mode, + target_partition_mode=preferred_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + if not current_to_preferred_converter.conversion_needed: + return hidden_states, padding_mask, input_ids + + current_to_preferred_converter.assert_no_dense_attention_inputs( + attention_mask=attention_mask, + hidden_states=hidden_states, + ) + hidden_states = current_to_preferred_converter.convert( + hidden_states, + seq_dim=0, + sequence_parallel=self.config.sequence_parallel, + ) + # Model-level rotary_pos_emb is only consumed by regular attention, whose + # CP layout preference is zigzag. MLA/CSA/DSv4-style variants must ignore + # external RoPE and manage any RoPE positions internally, so this layout + # edge intentionally does not convert rotary_pos_emb. + if padding_mask is not None: + padding_mask = current_to_preferred_converter.convert( + padding_mask, + seq_dim=1, + sequence_parallel=self.config.sequence_parallel, + ) + if input_ids is not None: + input_ids = current_to_preferred_converter.convert( + input_ids, + seq_dim=1, + ) + + return hidden_states, padding_mask, input_ids + def set_input_tensor(self, input_tensor: Tensor): """Set input tensor to be used instead of forward()'s input. @@ -931,6 +1008,23 @@ def get_inner_quant_context(config, layer_number): ) with outer_fp8_context: + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + cp_layout_needed = ( + cp_group is not None + and cp_group.size() > 1 + and self.config.cp_partition_mode == "auto" + ) + current_partition_mode = None + if cp_layout_needed: + current_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + self.cp_stage_entry_partition_mode, + owner_name=type(self).__name__, + cp_group=cp_group, + ) + packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, current_partition_mode + ) if self.config.recompute_granularity == 'full' and self.training: hidden_states = checkpointed_forward( self, @@ -947,6 +1041,32 @@ def get_inner_quant_context(config, layer_number): ) else: for l_no, layer in enumerate(self.layers): + if cp_layout_needed: + preferred_partition_mode = get_preferred_cp_partition_mode_for_layer( + layer, getattr(layer, "config", self.config) + ) + (hidden_states, padding_mask, input_ids) = ( + self._convert_cp_partition_mode_for_layer( + local_index=l_no, + current_partition_mode=current_partition_mode, + hidden_states=hidden_states, + attention_mask=attention_mask, + packed_seq_params=packed_seq_params, + padding_mask=padding_mask, + input_ids=input_ids, + preferred_partition_mode=preferred_partition_mode, + ) + ) + if preferred_partition_mode is not None: + packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, preferred_partition_mode + ) + current_partition_mode = getattr( + packed_seq_params, + "cp_partition_mode", + preferred_partition_mode or current_partition_mode, + ) + # Layers have 1-indexed layer numbers attribute. inner_quant_context = get_inner_quant_context( self.config, layer.layer_number - 1 diff --git a/megatron/core/models/hybrid/hybrid_layer_allocation.py b/megatron/core/models/hybrid/hybrid_layer_allocation.py index a8d2006c3b0..7f186348208 100644 --- a/megatron/core/models/hybrid/hybrid_layer_allocation.py +++ b/megatron/core/models/hybrid/hybrid_layer_allocation.py @@ -6,7 +6,13 @@ import torch -from megatron.core.utils import log_on_each_pipeline_stage, log_single_rank +from megatron.core.context_parallel_layout import CpPartitionMode +from megatron.core.utils import ( + get_pg_rank, + get_pg_size, + log_on_each_pipeline_stage, + log_single_rank, +) logger = logging.getLogger(__name__) @@ -201,6 +207,171 @@ def get_hybrid_layer_counts(pattern: str) -> Dict[str, int]: return counts +def get_hybrid_layer_cp_partition_mode(layer_symbol: str, config) -> Optional[CpPartitionMode]: + """Return the CP partition mode preferred by one hybrid layer symbol. + + ``None`` means the layer is token-layout agnostic and can preserve whatever + CP partition mode it receives. + """ + if layer_symbol == Symbols.MAMBA: + # MambaContextParallel currently undoes/redoes Megatron's attention + # load-balancing layout internally, so it expects zigzag inputs. + return "zigzag" + if layer_symbol == Symbols.GDN: + mode = getattr(config, "linear_cp_mode", "chunkwise") + if mode == "chunkwise": + return "contiguous" + if mode == "headwise": + return "zigzag" + raise ValueError(f"Unsupported GatedDeltaNet linear_cp_mode: {mode!r}.") + if layer_symbol == Symbols.ATTENTION: + return "zigzag" + if layer_symbol == Symbols.DS_ATTENTION: + if getattr(config, "experimental_attention_variant", None) == "dsv4_hybrid": + return "contiguous" + return "zigzag" + if layer_symbol in {Symbols.CSA, Symbols.HCA, Symbols.WINDOW}: + return "contiguous" + if layer_symbol in {Symbols.MLP, Symbols.MOE}: + return None + raise ValueError(f"Unsupported hybrid layer symbol {layer_symbol!r}.") + + +def get_hybrid_stage_input_cp_partition_mode( + config, hybrid_layer_pattern: Optional[str], stage_layer_offset: int +) -> Optional[CpPartitionMode]: + """Return the CP partition mode expected at one HybridModel stage input.""" + parsed = parse_hybrid_pattern(hybrid_layer_pattern) + main_pattern = (parsed.main_pattern or "").replace(Symbols.PIPE, "") + layer_layouts = [ + get_hybrid_layer_cp_partition_mode(layer_symbol, config) for layer_symbol in main_pattern + ] + + current_partition_mode = None + for preferred_partition_mode in layer_layouts[:stage_layer_offset]: + if preferred_partition_mode is not None: + current_partition_mode = preferred_partition_mode + + if current_partition_mode is None: + for preferred_partition_mode in layer_layouts[stage_layer_offset:]: + if preferred_partition_mode is not None: + current_partition_mode = preferred_partition_mode + break + + return current_partition_mode + + +def get_hybrid_stage_input_cp_partition_mode_for_stage( + config, + hybrid_layer_pattern: Optional[str], + pp_group: Optional[torch.distributed.ProcessGroup], + vp_stage: Optional[int], + *, + first_stage_layers: Optional[int] = None, + last_stage_layers: Optional[int] = None, +) -> Optional[CpPartitionMode]: + """Return the CP partition mode expected at one hybrid PP/VPP stage input.""" + parsed = parse_hybrid_pattern(hybrid_layer_pattern) + layer_offset = get_hybrid_stage_layer_offset( + parsed.main_pattern or '', + pp_group, + vp_stage, + first_stage_layers=first_stage_layers, + last_stage_layers=last_stage_layers, + ) + return get_hybrid_stage_input_cp_partition_mode(config, hybrid_layer_pattern, layer_offset) + + +def get_hybrid_stage_layer_offset( + main_pattern: str, + pp_group: Optional[torch.distributed.ProcessGroup], + vp_stage: Optional[int], + *, + first_stage_layers: Optional[int] = None, + last_stage_layers: Optional[int] = None, +) -> int: + """Return the global main-layer offset for one hybrid PP/VPP stage without logging.""" + segments = main_pattern.split(Symbols.PIPE) if main_pattern else [''] + + pp_rank = get_pg_rank(pp_group) + pp_size = get_pg_size(pp_group) + + if len(segments) > 1 and (first_stage_layers is not None or last_stage_layers is not None): + raise ValueError( + "Cannot specify num_layers_in_first_pipeline_stage or " + "num_layers_in_last_pipeline_stage when hybrid_layer_pattern " + "contains pipe ('|') separators. The pipeline layout is already " + "explicitly defined by the pipe separators." + ) + + if len(segments) == 1 and pp_size > 1: + if vp_stage is not None: + raise ValueError( + "Virtual pipeline parallelism (vp_stage != None) is not supported " + "when hybrid_layer_pattern has no pipe ('|') separators. " + "Add '|' separators to define explicit pipeline/virtual-pipeline " + "stage boundaries." + ) + layer_type_list = validate_segment_layers(segments[0]) + num_layers = len(layer_type_list) + + if first_stage_layers is not None or last_stage_layers is not None: + first = first_stage_layers or 0 + last = last_stage_layers or 0 + middle_num_layers = num_layers - first - last + middle_stages = pp_size - sum( + 1 for x in (first_stage_layers, last_stage_layers) if x is not None + ) + if middle_stages > 0: + if middle_num_layers % middle_stages != 0: + raise ValueError( + f"Middle layers ({middle_num_layers}) must be evenly divisible " + f"by middle pipeline stages ({middle_stages})." + ) + layers_per_middle = middle_num_layers // middle_stages + else: + layers_per_middle = 0 + + is_first = first_stage_layers is not None and pp_rank == 0 + is_last = last_stage_layers is not None and pp_rank == pp_size - 1 + + if is_first: + return 0 + if is_last: + return num_layers - last + + middle_rank = pp_rank if first_stage_layers is None else pp_rank - 1 + return middle_rank * layers_per_middle + first + + if num_layers % pp_size != 0: + raise ValueError( + f"Number of layers ({num_layers}) must be evenly divisible " + f"by pipeline-model-parallel-size ({pp_size}) when no pipe " + f"separators are specified in the pattern." + ) + return pp_rank * (num_layers // pp_size) + + if len(segments) > 1 and len(segments) % pp_size != 0: + raise ValueError( + f"The number of pipe-delimited segments ({len(segments)}) in " + f"hybrid_layer_pattern must be evenly divisible by " + f"pipeline_model_parallel_size ({pp_size})." + ) + + vp_rel = vp_stage if vp_stage is not None else 0 + segment_index = vp_rel * pp_size + pp_rank + if segment_index >= len(segments): + raise ValueError( + f"Pipeline segment index {segment_index} (pp_rank={pp_rank}, " + f"vp_stage={vp_rel}) is out of range for {len(segments)} segments. " + f"The pattern does not define enough pipe-delimited segments for " + f"the current PP/VPP configuration." + ) + + validate_segment_layers(segments[segment_index]) + return sum(len(segments[i]) for i in range(segment_index)) + + def parse_hybrid_pattern(pattern: Optional[str]) -> ParsedHybridPattern: """Parse a unified hybrid pattern string into main and MTP components. @@ -380,8 +551,8 @@ def select_pipeline_segment( """ segments = main_pattern.split(Symbols.PIPE) if main_pattern else [''] - pp_rank = torch.distributed.get_rank(pp_group) if pp_group is not None else 0 - pp_size = torch.distributed.get_world_size(pp_group) if pp_group is not None else 1 + pp_rank = get_pg_rank(pp_group) + pp_size = get_pg_size(pp_group) if len(segments) > 1 and (first_stage_layers is not None or last_stage_layers is not None): raise ValueError( diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 07c86880f9f..b17d353f187 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -1,19 +1,25 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging +import warnings from typing import Literal, Optional from torch import Tensor from megatron.core import tensor_parallel from megatron.core.config_logger import has_config_logger_enabled, log_config_to_disk +from megatron.core.context_parallel_layout import ( + CpPartitionModeConverter, + get_stage_entry_partition_mode, + replace_packed_seq_params_cp_partition_mode, +) from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.inference.utils import InferenceMode from megatron.core.models.common.embeddings.language_model_embedding import LanguageModelEmbedding from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding 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.packed_seq_params import PackedSeqParams +from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( FineGrainedActivationOffloadingInterface as off_interface, ) @@ -98,6 +104,8 @@ class HybridModel(LanguageModule, GraphableMegatronModule): Defaults to None. pg_collection (ProcessGroupCollection, optional): Model communication process groups. vp_stage (Optional[int], optional): Virtual pipeline stage index. Defaults to None. + cp_stage_entry_partition_mode (str, optional): CP partition mode expected at this + stage input. Required when context parallelism is enabled. Defaults to None. """ def __init__( @@ -115,6 +123,8 @@ def __init__( fp16_lm_cross_entropy: bool = False, parallel_output: bool = True, share_embeddings_and_output_weights: bool = False, + # TODO(yuzhongw): re-audit Mamba-specific comments in Hybrid + # model paths and keep only comments that are truly Mamba-layer specific. # Mamba with no attention has no need for position embeddings, so none is default position_embedding_type: Literal['learned_absolute', 'rope', 'yarn', 'none'] = 'none', rotary_percent: float = 1.0, @@ -123,6 +133,7 @@ def __init__( seq_len_interpolation_factor: Optional[float] = None, pg_collection: Optional[ProcessGroupCollection] = None, vp_stage: Optional[int] = None, + cp_stage_entry_partition_mode: Optional[str] = None, ) -> None: super().__init__(config=config, pg_collection=pg_collection) @@ -210,7 +221,6 @@ def __init__( last_stage_layers=self.config.num_layers_in_last_pipeline_stage, **logging_pg_kwargs, ) - # Determine if MTP is needed (based on pattern parsing) self.mtp_process = ( self.mtp_pattern is not None @@ -282,6 +292,7 @@ def __init__( dtype=config.params_dtype, pg_collection=self.pg_collection, name="decoder", + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, ) # MTP block - uses mtp_block_spec from hybrid_stack_spec.submodules @@ -444,6 +455,27 @@ def forward( inference_context = deprecate_inference_params(inference_context, inference_params) + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + if cp_group is not None and cp_group.size() > 1 and packed_seq_params is None: + warnings.warn( + "HybridModel received no PackedSeqParams while running under context " + "parallelism. Megatron-LM will temporarily assume SBHD tensors and " + "create layout metadata for this forward pass. In a future release, " + "callers must pass PackedSeqParams with qkv_format and cp_partition_mode " + "set explicitly.", + FutureWarning, + stacklevel=2, + ) + packed_seq_params = PackedSeqParams( + qkv_format="sbhd", cp_partition_mode=self.decoder.cp_stage_entry_partition_mode + ) + input_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + self.decoder.cp_stage_entry_partition_mode, + owner_name=type(self).__name__, + cp_group=cp_group, + ) + in_inference_mode = InferenceMode.is_active() if in_inference_mode: @@ -477,6 +509,9 @@ def forward( decoder_input = None rotary_pos_emb = None + # Model-level rotary_pos_emb is only for regular attention. Regular + # attention uses the default zigzag CP RoPE layout; MLA/CSA/DSv4-style + # variants must ignore this external RoPE and build/apply RoPE internally. if 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 @@ -484,6 +519,7 @@ def forward( rotary_pos_emb = self.rotary_pos_emb( rotary_seq_len, packed_seq=packed_seq_params is not None and packed_seq_params.qkv_format == 'thd', + cp_group=packed_seq_params.cp_group if packed_seq_params is not None else None, ) elif self.position_embedding_type == 'yarn': rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( @@ -493,6 +529,7 @@ def forward( rotary_pos_emb, _ = self.rotary_pos_emb( rotary_seq_len, packed_seq=packed_seq_params is not None and packed_seq_params.qkv_format == 'thd', + cp_group=packed_seq_params.cp_group if packed_seq_params is not None else None, ) # Wrap decoder_input to allow the decoder (HybridStack) to delete the @@ -538,10 +575,6 @@ def forward( hidden_states = decoder_output mhc_multistream = None - output_weight = None - if self.share_embeddings_and_output_weights: - output_weight = self.shared_embedding_or_output_weight() - # Check if speculative decoding is active. When it is, MTP must be # computed *after* verification so that it is conditioned on verified # tokens rather than stale speculative tokens from the previous step. @@ -552,7 +585,77 @@ def forward( and inference_context.num_speculative_tokens > 0 ) + output_weight = None + if self.share_embeddings_and_output_weights: + output_weight = self.shared_embedding_or_output_weight() + + postprocess_to_input_converter = None mtp_forward_ran = self.mtp_process and not (in_inference_mode or is_spec_decode) + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + cp_size = cp_group.size() if cp_group is not None else 1 + + needs_batch_layout = ( + cp_size > 1 + and self.config.cp_partition_mode == "auto" + and (self.post_process or mtp_forward_ran) + ) + if needs_batch_layout: + block_output_partition_mode = getattr( + packed_seq_params, "cp_partition_mode", input_partition_mode + ) + postprocess_partition_mode = ( + block_output_partition_mode if mtp_forward_ran else input_partition_mode + ) + block_to_postprocess_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=block_output_partition_mode, + target_partition_mode=postprocess_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + hidden_states = block_to_postprocess_converter.convert( + hidden_states, seq_dim=0, sequence_parallel=self.config.sequence_parallel + ) + if mhc_multistream is not None: + mhc_multistream = block_to_postprocess_converter.convert( + mhc_multistream, seq_dim=0, sequence_parallel=self.config.sequence_parallel + ) + if input_partition_mode != postprocess_partition_mode: + input_to_postprocess_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=input_partition_mode, + target_partition_mode=postprocess_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + input_to_postprocess_converter.assert_no_dense_attention_inputs( + attention_mask=attention_mask, + hidden_states=hidden_states, + ) + input_ids = input_to_postprocess_converter.convert(input_ids, seq_dim=-1) + position_ids = input_to_postprocess_converter.convert(position_ids, seq_dim=-1) + labels = input_to_postprocess_converter.convert(labels, seq_dim=-1) + loss_mask = input_to_postprocess_converter.convert(loss_mask, seq_dim=-1) + padding_mask = input_to_postprocess_converter.convert( + padding_mask, seq_dim=-1, sequence_parallel=self.config.sequence_parallel + ) + # Model-level rotary_pos_emb belongs to regular attention, whose + # CP layout preference is zigzag. MTP side tensors are aligned + # for token/loss semantics, but RoPE is not treated as a batch + # side tensor to be converted here. + postprocess_to_input_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=postprocess_partition_mode, + target_partition_mode=input_partition_mode, + config=self.config, + ) + packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, postprocess_partition_mode + ) + if mtp_forward_ran: hidden_states = self.mtp( input_ids=input_ids, @@ -577,6 +680,7 @@ def forward( else: # For RL (labels is None), process_mtp_loss derives labels from # input_ids to match the SFT label format. + mtp_cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) hidden_states = process_mtp_loss( hidden_states=hidden_states, labels=labels, @@ -587,7 +691,7 @@ def forward( is_training=self.training, compute_language_model_loss=self.compute_language_model_loss, config=self.config, - cp_group=self.pg_collection.cp, + cp_group=mtp_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, @@ -633,8 +737,13 @@ def forward( if labels is None: # [s b h] => [b s h] + if postprocess_to_input_converter is not None: + logits = postprocess_to_input_converter.convert(logits, seq_dim=0) return logits.transpose(0, 1).contiguous() loss = self.compute_language_model_loss(labels, logits) + if postprocess_to_input_converter is not None: + loss = postprocess_to_input_converter.convert(loss, seq_dim=-1) + return loss diff --git a/megatron/core/models/multimodal/llava_model.py b/megatron/core/models/multimodal/llava_model.py index f714afeca64..994f896011b 100644 --- a/megatron/core/models/multimodal/llava_model.py +++ b/megatron/core/models/multimodal/llava_model.py @@ -183,6 +183,11 @@ def __init__( self.pg_collection = pg_collection language_model_type = getattr(language_transformer_config, "language_model_type", "") + if getattr(language_transformer_config, "cp_partition_mode", "zigzag") == "auto": + raise ValueError( + 'LLaVAModel does not support cp_partition_mode="auto"; use a pretraining ' + 'entrypoint that owns CP layout prebuild and batch partitioning.' + ) self.sequence_parallel_lm = language_transformer_config.sequence_parallel self.tp_comm_overlap_lm = language_transformer_config.tp_comm_overlap self.context_parallel_lm = language_transformer_config.context_parallel_size @@ -871,7 +876,9 @@ def _process_embedding_token_parallel( from megatron.core.utils import get_batch_on_this_cp_rank batch = get_batch_on_this_cp_rank( - batch, is_hybrid_cp=False, cp_group=get_context_parallel_group() + batch, + is_hybrid_cp=False, + cp_group=get_context_parallel_group(), ) else: assert HAVE_TEX and is_te_min_version( diff --git a/megatron/core/packed_seq_params.py b/megatron/core/packed_seq_params.py index a0957a39eab..204a76f4661 100644 --- a/megatron/core/packed_seq_params.py +++ b/megatron/core/packed_seq_params.py @@ -7,12 +7,21 @@ import torch.nn.functional as F from torch import Tensor +THD_CP_PARTITION_ROUTE_TENSOR_FIELDS = ( + "cp_partition_route_zigzag_to_contiguous", + "cp_partition_route_contiguous_to_zigzag", +) + @dataclass class PackedSeqParams: ''' parameters to TEDotProductAttention and fused rope kernels for the `thd` (packed) sequence format + + ``cp_partition_route_*`` tensors are per-microbatch THD CP layout + conversion routes. Metadata annotation helpers update the current + partition mode in-place while preserving these route tensor identities. ''' qkv_format: str = None @@ -27,8 +36,10 @@ class PackedSeqParams: total_tokens: int = None seq_idx: Tensor = None pad_between_seqs: Optional[bool] = None - cp_partition_mode: Literal["zigzag", "contiguous"] = "zigzag" + cp_partition_mode: Optional[Literal["zigzag", "contiguous"]] = None tokens_per_sample: int = None + cp_partition_route_zigzag_to_contiguous: Tensor = None + cp_partition_route_contiguous_to_zigzag: Tensor = None def __post_init__(self): """Pre-compute seq_idx for Mamba mixer CUDA graph compatibility. @@ -43,6 +54,15 @@ def __post_init__(self): cu_seqlens_q_padded[-1] == max_seqlen then this additional sequence index will not be included. """ + if self.cp_partition_mode is not None and self.cp_partition_mode not in ( + "zigzag", + "contiguous", + ): + raise ValueError( + "PackedSeqParams.cp_partition_mode must be a concrete runtime layout " + f"('zigzag' or 'contiguous'), got {self.cp_partition_mode!r}." + ) + cu_seqlens = ( self.cu_seqlens_q_padded if self.cu_seqlens_q_padded is not None else self.cu_seqlens_q ) diff --git a/megatron/core/recompute.py b/megatron/core/recompute.py index 8974efc1311..922aab47f61 100644 --- a/megatron/core/recompute.py +++ b/megatron/core/recompute.py @@ -5,10 +5,16 @@ from torch import Tensor from megatron.core import tensor_parallel +from megatron.core.context_parallel_layout import ( + CpPartitionModeConverter, + get_preferred_cp_partition_mode_for_layer, + get_stage_entry_partition_mode, + replace_packed_seq_params_cp_partition_mode, +) from megatron.core.extensions.transformer_engine import HAVE_TE from megatron.core.fp4_utils import get_fp4_context from megatron.core.fp8_utils import get_fp8_context -from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_layer import TransformerLayer @@ -50,15 +56,101 @@ def checkpointed_forward( if extract_layer_indices is None: extract_layer_indices = set() intermediate_hidden_states: List[Tensor] = [] + cp_group = resolve_cp_group(getattr(self.pg_collection, "cp", None), packed_seq_params) + cp_layout_needed = ( + cp_group is not None + and cp_group.size() > 1 + and self.config.cp_partition_mode == "auto" + ) + stage_entry_partition_mode = ( + get_stage_entry_partition_mode( + packed_seq_params, + getattr(self, "cp_stage_entry_partition_mode", None), + owner_name=type(self).__name__, + cp_group=cp_group, + ) + if cp_layout_needed + else None + ) def custom(start: int, end: int): def custom_forward( hidden_states, attention_mask, context, context_mask, rotary_pos_emb, padding_mask=None ): + current_partition_mode = stage_entry_partition_mode + if current_partition_mode is not None: + for index in range(start): + # Use self.layers[index] (not self._get_layer) so this + # function works for both TransformerBlock and HybridStack. + layer = self.layers[index] + preferred_partition_mode = get_preferred_cp_partition_mode_for_layer( + layer, getattr(layer, "config", self.config) + ) + if preferred_partition_mode is not None: + current_partition_mode = preferred_partition_mode + local_packed_seq_params = packed_seq_params + local_input_ids = input_ids + if current_partition_mode is not None: + local_packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, current_partition_mode + ) + if cp_layout_needed and current_partition_mode != stage_entry_partition_mode: + chunk_entry_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=local_packed_seq_params, + source_partition_mode=stage_entry_partition_mode, + target_partition_mode=current_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + chunk_entry_converter.assert_no_dense_attention_inputs( + attention_mask=attention_mask, + attention_bias=attention_bias, + hidden_states=hidden_states, + ) + # Checkpoint chunks may enter after earlier layers changed the + # current CP layout. Re-align batch/token side tensors here, but + # keep model-level RoPE in the regular-attention zigzag layout. + if padding_mask is not None: + padding_mask = chunk_entry_converter.convert( + padding_mask, + seq_dim=1, + sequence_parallel=self.config.sequence_parallel, + ) + if local_input_ids is not None: + local_input_ids = chunk_entry_converter.convert( + local_input_ids, + seq_dim=1, + ) for index in range(start, end): # Use self.layers[index] (not self._get_layer) so this # function works for both TransformerBlock and HybridStack. layer = self.layers[index] + if cp_layout_needed: + preferred_partition_mode = get_preferred_cp_partition_mode_for_layer( + layer, getattr(layer, "config", self.config) + ) + (hidden_states, padding_mask, local_input_ids) = ( + self._convert_cp_partition_mode_for_layer( + local_index=index, + current_partition_mode=current_partition_mode, + hidden_states=hidden_states, + attention_mask=attention_mask, + packed_seq_params=local_packed_seq_params, + padding_mask=padding_mask, + input_ids=local_input_ids, + preferred_partition_mode=preferred_partition_mode, + ) + ) + if preferred_partition_mode is not None: + local_packed_seq_params = replace_packed_seq_params_cp_partition_mode( + local_packed_seq_params, preferred_partition_mode + ) + current_partition_mode = getattr( + local_packed_seq_params, + "cp_partition_mode", + preferred_partition_mode or current_partition_mode, + ) # Get appropriate inner quantization context if use_inner_quantization_context: @@ -87,13 +179,17 @@ def custom_forward( rotary_pos_emb=rotary_pos_emb, attention_bias=attention_bias, inference_context=None, - packed_seq_params=packed_seq_params, + packed_seq_params=local_packed_seq_params, padding_mask=padding_mask, - input_ids=input_ids, + input_ids=local_input_ids, ) with inner_quantization_context: if isinstance(layer, TransformerLayer): hidden_states, context = layer(**layer_kwargs) + elif layer.__class__.__name__ == "HyperConnectionHybridLayer": + for k in ("context", "context_mask", "attention_bias"): + layer_kwargs.pop(k, None) + hidden_states, context = layer(**layer_kwargs) else: # MambaLayer (HybridStack `M` slot) for k in ( "context", diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net.py index 7e28691c15c..66375be3f1b 100644 --- a/megatron/core/ssm/gated_delta_net.py +++ b/megatron/core/ssm/gated_delta_net.py @@ -6,6 +6,7 @@ # LICENSE file in the root directory of this source tree. import logging +import math from dataclasses import dataclass from functools import lru_cache from typing import Optional, Union @@ -16,10 +17,7 @@ from torch import Tensor from megatron.core import tensor_parallel -from megatron.core.context_parallel_layout import ( - contiguous_to_zigzag_chunks, - zigzag_to_contiguous_chunks, -) +from megatron.core.context_parallel_layout import convert_module_input_tensors_cp_partition_mode from megatron.core.fp8_utils import get_fp8_align_size from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.jit import jit_fuser @@ -80,6 +78,15 @@ class GatedDeltaNet(MegatronModule): and returns output of the same size. """ + def get_preferred_cp_partition_mode(self): + """Return GDN's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + mode = getattr(self.config, "linear_cp_mode", "chunkwise") + if mode == "chunkwise": + return "contiguous" + if mode == "headwise": + return "zigzag" + raise ValueError(f"Unsupported GatedDeltaNet linear_cp_mode: {mode!r}.") + def __init__( self, config: TransformerConfig, @@ -202,16 +209,19 @@ def __init__( # weight shape: [conv_dim, 1, d_conv] # bias shape: [conv_dim] - self.conv1d = nn.Conv1d( - in_channels=self.conv_dim_local_tp, - out_channels=self.conv_dim_local_tp, - bias=conv_bias, - kernel_size=self.conv_kernel_dim, - groups=self.conv_dim_local_tp, - padding=self.conv_kernel_dim - 1, - device=torch.cuda.current_device(), - dtype=config.params_dtype, - ) + # Conv1d performs implicit CUDA RNG initialization in its constructor. + # Isolate it; reset_parameters below owns the final Megatron-tracked init. + with torch.random.fork_rng(devices=[torch.cuda.current_device()]): + self.conv1d = nn.Conv1d( + in_channels=self.conv_dim_local_tp, + out_channels=self.conv_dim_local_tp, + bias=conv_bias, + kernel_size=self.conv_kernel_dim, + groups=self.conv_dim_local_tp, + padding=self.conv_kernel_dim - 1, + device=torch.cuda.current_device(), + dtype=config.params_dtype, + ) setattr(self.conv1d.weight, "tensor_model_parallel", True) setattr(self.conv1d.weight, "partition_dim", 0) if conv_bias: @@ -295,6 +305,12 @@ def reset_parameters(self): # conv1d.weight if self.conv_init is not None: nn.init.uniform_(self.conv1d.weight, -self.conv_init, self.conv_init) + else: + nn.init.kaiming_uniform_(self.conv1d.weight, a=math.sqrt(5)) + if self.conv1d.bias is not None: + fan_in = self.conv1d.weight.size(1) * self.conv1d.weight.size(2) + bound = 1 / math.sqrt(fan_in) + nn.init.uniform_(self.conv1d.bias, -bound, bound) # dt_bias torch.ones( self.num_v_heads_local_tp, @@ -366,6 +382,21 @@ def forward( ) cp_size_chunkwise = cp_group_chunkwise.size() if cp_group_chunkwise is not None else 1 cp_size_headwise = cp_group_headwise.size() if cp_group_headwise is not None else 1 + back_to_input_converter = None + if self.config.linear_cp_mode == "chunkwise": + ( + hidden_states, + packed_seq_params, + back_to_input_converter, + ) = convert_module_input_tensors_cp_partition_mode( + hidden_states=hidden_states, + packed_seq_params=packed_seq_params, + cp_group=cp_group_chunkwise, + tp_group=self.tp_group, + target_partition_mode="contiguous", + sequence_parallel=self.config.sequence_parallel, + config=self.config, + ) seq_len_local, batch, _ = hidden_states.shape seq_len_post_headwise = seq_len_local * self.sp_size * cp_size_headwise @@ -379,6 +410,29 @@ def forward( # TODO: support inference raise NotImplementedError("GDN does not support inference for now.") + if cp_size_chunkwise > 1: + expected_cp_partition_mode = "contiguous" + elif cp_size_headwise > 1: + expected_cp_partition_mode = "zigzag" + else: + expected_cp_partition_mode = None + actual_cp_partition_mode = getattr(packed_seq_params, "cp_partition_mode", None) + if expected_cp_partition_mode is not None and actual_cp_partition_mode is None: + raise ValueError( + "GatedDeltaNet requires PackedSeqParams.cp_partition_mode when context " + "parallelism is active." + ) + if ( + expected_cp_partition_mode is not None + and actual_cp_partition_mode != expected_cp_partition_mode + ): + raise ValueError( + f"GatedDeltaNet with linear_cp_mode={self.config.linear_cp_mode!r} prefers " + f"cp_partition_mode={expected_cp_partition_mode!r}, but packed_seq_params " + f"has {actual_cp_partition_mode!r}. CP partition conversion must be handled " + "before calling GatedDeltaNet." + ) + if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': assert batch == 1, "Packed sequence expects batch dimension to be 1" assert ( @@ -393,14 +447,14 @@ def forward( packed_seq_params.cu_seqlens_q, seq_len_global, "cu_seqlens_q", - cp_size=self.cp_size, + cp_size=cp_size_chunkwise, ) cu_seqlens_kv = self._resolve_cu_seqlens( packed_seq_params.cu_seqlens_kv_padded, packed_seq_params.cu_seqlens_kv, seq_len_global, "cu_seqlens_kv", - cp_size=self.cp_size, + cp_size=cp_size_chunkwise, ) assert torch.equal(cu_seqlens_q, cu_seqlens_kv), ( "Currently only support cu_seqlens_q equals to cu_seqlens_kv, " @@ -479,6 +533,13 @@ def _checkpointed_compute(hidden_states): chunkwise_cp_context, ) + if back_to_input_converter is not None: + out = back_to_input_converter.convert( + out, + seq_dim=0, + sequence_parallel=self.config.sequence_parallel, + ) + return out, out_bias def _forward_compute( @@ -508,23 +569,6 @@ def _forward_compute( qkvzba, _ = self.in_proj(hidden_states) nvtx_range_pop(suffix="in_proj") - # Chunkwise CP expects the contiguous-time chunk layout (rank r holds chunks - # [2r, 2r+1]) inside conv1d / chunk_gated_delta_rule. Megatron attention CP - # feeds us the zigzag attention-load-balanced layout (rank r holds - # [r, 2*cp-r-1]), so reshuffle chunks over the CP group with a single - # all-to-all — no full-sequence gather required. - # TODO: Move CP layout ownership to a model/region-level scheduler so hybrid models can - # enter contiguous layout before GDN regions instead of paying module-local conversions. - if cp_size_chunkwise > 1: - nvtx_range_push(suffix="zigzag_to_contiguous") - if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': - qkvzba = zigzag_to_contiguous_chunks( - qkvzba, cp_group_chunkwise, seq_dim=0, cu_seqlens=cu_seqlens_q - ) - else: - qkvzba = zigzag_to_contiguous_chunks(qkvzba, cp_group_chunkwise, seq_dim=0) - nvtx_range_pop(suffix="zigzag_to_contiguous") - qkvzba, thd_cp_a2a_inv = self._a2a_cp_to_hp( qkvzba, cp_size_headwise, @@ -621,23 +665,6 @@ def _forward_compute( norm_out = norm_out.reshape(batch, seq_len_post_headwise, -1) norm_out = norm_out.transpose(0, 1).contiguous() - # Inverse of the zigzag -> contiguous reshuffle performed before conv1d. - # Restores the Megatron attention-load-balanced layout that downstream - # layers and loss computation expect. - # TODO: The planned CP layout refactor should keep consecutive GDN layers contiguous and - # restore zigzag only at SDPA/canonical-layout boundaries. - if cp_size_chunkwise > 1: - nvtx_range_push(suffix="contiguous_to_zigzag") - if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': - norm_out = contiguous_to_zigzag_chunks( - norm_out, cp_group=cp_group_chunkwise, seq_dim=0, cu_seqlens=cu_seqlens_q - ) - else: - norm_out = contiguous_to_zigzag_chunks( - norm_out, cp_group=cp_group_chunkwise, seq_dim=0 - ) - nvtx_range_pop(suffix="contiguous_to_zigzag") - norm_out = self._a2a_hp_to_cp( norm_out, cp_size_headwise, cp_group_headwise, packed_seq_params, thd_cp_a2a_inv ) diff --git a/megatron/core/ssm/mamba_layer.py b/megatron/core/ssm/mamba_layer.py index d3b04e59c29..835439b5509 100644 --- a/megatron/core/ssm/mamba_layer.py +++ b/megatron/core/ssm/mamba_layer.py @@ -65,6 +65,10 @@ class MambaLayer(GraphableMegatronModule): output of the same size. """ + def get_preferred_cp_partition_mode(self): + """Return Mamba's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return "zigzag" + def __init__( self, config: TransformerConfig, diff --git a/megatron/core/ssm/mlp_layer.py b/megatron/core/ssm/mlp_layer.py index 14500e5ad18..fa235fb43c8 100644 --- a/megatron/core/ssm/mlp_layer.py +++ b/megatron/core/ssm/mlp_layer.py @@ -13,6 +13,10 @@ class MLPLayer(TransformerLayer): """Drop-in replacement for TransformerLayer but initializes only an MLP via the spec.""" + def get_preferred_cp_partition_mode(self): + """Return MLPLayer's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return None + def __init__( self, config: TransformerConfig, diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index 1f29c93eef3..a3e75317846 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -11,6 +11,7 @@ from torch import Tensor from megatron.core import tensor_parallel +from megatron.core.context_parallel_layout import convert_module_input_tensors_cp_partition_mode from megatron.core.dist_checkpointing import ShardedTensor from megatron.core.dist_checkpointing.mapping import ( ReplicaId, @@ -313,6 +314,7 @@ def __init__( self.attn_mask_type = attn_mask_type self.attention_type = attention_type + self.cp_comm_type = cp_comm_type self.batch_invariant_mode = config.batch_invariant_mode # Cache the YaRN concentration factor (a.k.a. attention factor / mscale), @@ -444,6 +446,10 @@ def __init__( rotary_base = self.config.rotary_base_per_layer[self.layer_number - 1] self._build_per_layer_rotary_pos_emb(rotary_base) + def get_preferred_cp_partition_mode(self): + """Return Attention's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return "zigzag" + def _build_per_layer_rotary_pos_emb(self, rotary_base: float) -> None: """Build self.rotary_pos_emb using a layer-specific rotary base.""" from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding @@ -1361,6 +1367,42 @@ def forward( if packed_seq_params is not None and packed_seq_params.local_cp_size is not None: assert packed_seq_params.cp_group is not None, "cp_group must be set in dynamic-cp mode" self.pg_collection.cp = packed_seq_params.cp_group + ( + hidden_states, + packed_seq_params, + back_to_input_converter, + ) = convert_module_input_tensors_cp_partition_mode( + hidden_states=hidden_states, + key_value_states=key_value_states, + packed_seq_params=packed_seq_params, + cp_group=self.pg_collection.cp, + tp_group=self.pg_collection.tp, + target_partition_mode=self.get_preferred_cp_partition_mode(), + sequence_parallel=self.config.sequence_parallel, + config=self.config, + attention_mask=attention_mask, + attention_bias=attention_bias, + ) + preferred_cp_partition_mode = self.get_preferred_cp_partition_mode() + if packed_seq_params is not None and preferred_cp_partition_mode is not None: + cp_group = packed_seq_params.cp_group + if cp_group is None and hasattr(self.pg_collection, 'cp'): + cp_group = self.pg_collection.cp + if cp_group is not None and get_pg_size(cp_group) > 1: + actual_cp_partition_mode = packed_seq_params.cp_partition_mode + if actual_cp_partition_mode is None: + raise ValueError( + f"{self.__class__.__name__} requires " + "PackedSeqParams.cp_partition_mode when context parallelism is active." + ) + if actual_cp_partition_mode != preferred_cp_partition_mode: + raise ValueError( + f"{self.__class__.__name__} prefers " + f"cp_partition_mode={preferred_cp_partition_mode!r}, but " + f"packed_seq_params has {actual_cp_partition_mode!r}. CP partition " + "conversion must be handled by TransformerBlock before entering " + "attention." + ) # Check if we need to skip RoPE # no_rope is 0-indexed array and self.layer_number is 1-indexed @@ -1700,6 +1742,13 @@ def forward( output = attn_proj_manager.group_offload(output, forced_released_tensors=[core_attn_out]) nvtx_range_pop(suffix="linear_proj") + if back_to_input_converter is not None: + output = back_to_input_converter.convert( + output, + seq_dim=0, + sequence_parallel=self.config.sequence_parallel, + ) + self.pg_collection.cp = _orig_cp_group return output, bias diff --git a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py index 64a08cb7837..166e96fd373 100644 --- a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py +++ b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py @@ -142,6 +142,10 @@ class AbsorbedMLASelfAttention(Attention): computation which can be more efficient for certain attention variants. """ + def get_preferred_cp_partition_mode(self): + """Return AbsorbedMLA's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return "zigzag" + def __init__( self, config: MLATransformerConfig, @@ -916,6 +920,25 @@ def forward( if packed_seq_params is not None and packed_seq_params.local_cp_size is not None: assert packed_seq_params.cp_group is not None, "cp_group must be set in dynamic-cp mode" self.pg_collection.cp = packed_seq_params.cp_group + if packed_seq_params is not None: + cp_group = packed_seq_params.cp_group + if cp_group is None: + cp_group = self.pg_collection.cp + if cp_group is not None and cp_group.size() > 1: + preferred_cp_partition_mode = self.get_preferred_cp_partition_mode() + actual_cp_partition_mode = packed_seq_params.cp_partition_mode + if actual_cp_partition_mode is None: + raise ValueError( + "AbsorbedMLASelfAttention requires PackedSeqParams.cp_partition_mode " + "when context parallelism is active." + ) + if actual_cp_partition_mode != preferred_cp_partition_mode: + raise ValueError( + "AbsorbedMLASelfAttention prefers " + f"cp_partition_mode={preferred_cp_partition_mode!r}, but " + f"packed_seq_params has {actual_cp_partition_mode!r}. CP partition " + "conversion must be handled before entering AbsorbedMLA." + ) # ===================== # Query, Key, and Value diff --git a/megatron/core/transformer/experimental_attention_variant/csa.py b/megatron/core/transformer/experimental_attention_variant/csa.py index 834482a244f..b68763327e5 100644 --- a/megatron/core/transformer/experimental_attention_variant/csa.py +++ b/megatron/core/transformer/experimental_attention_variant/csa.py @@ -1706,6 +1706,10 @@ class CompressedSparseAttention(MegatronModule): * ``ratio == 128``: window + 128x compressed, attend to all (compressor built only) """ + def get_preferred_cp_partition_mode(self): + """Return CSA's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return "contiguous" + def __init__( self, config: TransformerConfig, @@ -2087,8 +2091,29 @@ def forward( """ nvtx_range_push("compressed_sparse_attn") + cp_group = None + if packed_seq_params is not None: + cp_group = packed_seq_params.cp_group + if cp_group is None: + cp_group = self.pg_collection.cp + if cp_group is not None and cp_group.size() > 1: + preferred_cp_partition_mode = self.get_preferred_cp_partition_mode() + actual_cp_partition_mode = packed_seq_params.cp_partition_mode + if actual_cp_partition_mode is None: + raise ValueError( + "CompressedSparseAttention requires PackedSeqParams.cp_partition_mode " + "when context parallelism is active." + ) + if actual_cp_partition_mode != preferred_cp_partition_mode: + raise ValueError( + "CompressedSparseAttention prefers " + f"cp_partition_mode={preferred_cp_partition_mode!r}, but " + f"packed_seq_params has {actual_cp_partition_mode!r}. CP partition " + "conversion must be handled before entering CSA." + ) + if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': - if self.pg_collection.cp is not None and self.pg_collection.cp.size() > 1: + if cp_group is not None and cp_group.size() > 1: output = self._forward_thd_cp( query, key, x, qr, boundary_hidden, boundary_kv, packed_seq_params ) diff --git a/megatron/core/transformer/experimental_attention_variant/deepseek_v4_hybrid_attention.py b/megatron/core/transformer/experimental_attention_variant/deepseek_v4_hybrid_attention.py index 3c8f139d8a0..ee337030df6 100644 --- a/megatron/core/transformer/experimental_attention_variant/deepseek_v4_hybrid_attention.py +++ b/megatron/core/transformer/experimental_attention_variant/deepseek_v4_hybrid_attention.py @@ -59,6 +59,10 @@ class DSv4HybridSelfAttentionSubmodules: class DSv4HybridAttention(Attention): """DeepSeek-v4 Hybrid Attention layer.""" + def get_preferred_cp_partition_mode(self): + """Return this variant's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return "contiguous" + def __init__( self, config: MLATransformerConfig, @@ -267,8 +271,21 @@ def forward( if cp_size > 1 and qkv_format != 'thd': raise ValueError("DSv4 Hybrid with CP requires qkv_format='thd'.") use_thd_cp = cp_size > 1 and qkv_format == 'thd' - if use_thd_cp and packed_seq_params.cp_partition_mode != "contiguous": - raise ValueError("DSv4 THD CP requires a contiguous CP partition.") + if use_thd_cp: + preferred_cp_partition_mode = self.get_preferred_cp_partition_mode() + actual_cp_partition_mode = packed_seq_params.cp_partition_mode + if actual_cp_partition_mode is None: + raise ValueError( + "DSv4HybridAttention requires PackedSeqParams.cp_partition_mode " + "when context parallelism is active." + ) + if actual_cp_partition_mode != preferred_cp_partition_mode: + raise ValueError( + "DSv4HybridAttention prefers " + f"cp_partition_mode={preferred_cp_partition_mode!r}, but " + f"packed_seq_params has {actual_cp_partition_mode!r}. CP partition " + "conversion must be handled before entering DSv4HybridAttention." + ) self.pg_collection.cp = cp_group boundary_hidden = None diff --git a/megatron/core/transformer/identity_op.py b/megatron/core/transformer/identity_op.py index 6d42beb5a8f..5c67febde0d 100644 --- a/megatron/core/transformer/identity_op.py +++ b/megatron/core/transformer/identity_op.py @@ -11,6 +11,10 @@ class IdentityOp(torch.nn.Module): This is a placeholder for IdentityOp(x) -> x """ + def get_preferred_cp_partition_mode(self): + """Return IdentityOp's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return None + def __init__(self, *args: object, **kwargs: object): super().__init__() diff --git a/megatron/core/transformer/multi_latent_attention.py b/megatron/core/transformer/multi_latent_attention.py index 439a0f0649e..6e88eb0a827 100644 --- a/megatron/core/transformer/multi_latent_attention.py +++ b/megatron/core/transformer/multi_latent_attention.py @@ -134,6 +134,10 @@ class MultiLatentAttention(Attention): "cross attn" specializations. """ + def get_preferred_cp_partition_mode(self): + """Return MLA's CP layout preference for ``cp_partition_mode="auto"`` rollout.""" + return "zigzag" + def __init__( self, config: MLATransformerConfig, @@ -349,6 +353,25 @@ def forward( if packed_seq_params is not None and packed_seq_params.local_cp_size is not None: assert packed_seq_params.cp_group is not None, "cp_group must be set in dynamic-cp mode" self.pg_collection.cp = packed_seq_params.cp_group + if packed_seq_params is not None: + cp_group = packed_seq_params.cp_group + if cp_group is None: + cp_group = self.pg_collection.cp + if cp_group is not None and cp_group.size() > 1: + preferred_cp_partition_mode = self.get_preferred_cp_partition_mode() + actual_cp_partition_mode = packed_seq_params.cp_partition_mode + if actual_cp_partition_mode is None: + raise ValueError( + "MultiLatentAttention requires PackedSeqParams.cp_partition_mode " + "when context parallelism is active." + ) + if actual_cp_partition_mode != preferred_cp_partition_mode: + raise ValueError( + "MultiLatentAttention prefers " + f"cp_partition_mode={preferred_cp_partition_mode!r}, but " + f"packed_seq_params has {actual_cp_partition_mode!r}. CP partition " + "conversion must be handled before entering MLA." + ) # ===================== # Query, Key, and Value diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 60b8867e496..d02ecc4c310 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -11,6 +11,7 @@ from torch import Tensor from megatron.core import InferenceParams, parallel_state, tensor_parallel +from megatron.core.context_parallel_layout import get_preferred_cp_partition_mode_for_layer from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import apply_prefix_mapping, replace_prefix_for_sharding from megatron.core.enums import Fp8Recipe @@ -156,8 +157,9 @@ def roll_tensor(tensor, shifts=-1, dims=-1, cp_group=None, packed_seq_params=Non dims (int): The dimension to roll (typically -1 for sequence dimension). cp_group (ProcessGroup): The context parallelism process group. If None or size=1, falls back to standard rolling behavior. - packed_seq_params (PackedSeqParams): Parameters for packed sequence processing. - If provided, respects sequence boundaries. + packed_seq_params (PackedSeqParams): Parameters for sequence metadata. With + ``qkv_format='thd'``, rolling respects packed sequence boundaries. + Under CP>1, it must provide the current ``cp_partition_mode``. fill_value: Value to fill at boundary positions where the original sequence has no data (default 0). For most tensors (input_ids, loss_mask, labels) zero is correct. For a padding_mask with True=padded convention, @@ -168,8 +170,8 @@ def roll_tensor(tensor, shifts=-1, dims=-1, cp_group=None, packed_seq_params=Non if tensor is None: return None, None - # Handle packed sequences cases - if packed_seq_params is not None: + # Packed THD needs sequence-boundary-aware rolling even without CP. + if getattr(packed_seq_params, 'qkv_format', None) == 'thd': return _roll_tensor_packed_seq( tensor, shifts, dims, packed_seq_params, cp_group, fill_value=fill_value ) @@ -180,6 +182,14 @@ def roll_tensor(tensor, shifts=-1, dims=-1, cp_group=None, packed_seq_params=Non rolled_tensor.select(dims, shifts).fill_(fill_value) return rolled_tensor, rolled_tensor.sum() + cp_partition_mode = getattr(packed_seq_params, 'cp_partition_mode', None) + if cp_partition_mode != "zigzag": + raise NotImplementedError( + "MTP rolling with non-packed CP currently supports only zigzag layout; " + f"got {cp_partition_mode!r}. Contiguous layout for non-packed/SBHD MTP is " + "not supported yet." + ) + # CP-enabled rolling: Split tensor into chunks and handle boundary communication # This matches the batch splitting logic in get_batch_on_this_cp_rank() function tensor_list = tensor.chunk(2, dim=dims) @@ -277,7 +287,12 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No rolled_tensor[..., start_idx:end_idx] = rolled_seq return rolled_tensor, rolled_tensor.sum() - cp_partition_mode = getattr(packed_seq_params, 'cp_partition_mode', 'zigzag') + cp_partition_mode = getattr(packed_seq_params, 'cp_partition_mode', None) + if cp_partition_mode is None: + raise ValueError( + "PackedSeqParams.cp_partition_mode must be set when rolling packed sequences " + "under context parallelism." + ) if cp_partition_mode == 'zigzag': rolled_tensor = _roll_tensor_packed_seq_zigzag_cp( tensor, shifts, dims, cu_seqlens, cp_group, fill_value=fill_value @@ -293,6 +308,10 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No def _roll_tensor_packed_seq_zigzag_cp(tensor, shifts, dims, cu_seqlens, cp_group, fill_value=0): """Roll a zigzag-CP THD shard without crossing packed sequence boundaries.""" + # TODO(yuzhongw): replace the per-sequence boundary exchange + # with a route/boundary helper that batches communication across packed + # sequences. Keep this path as the reference until the optimized path has + # randomized packed-sequence parity coverage. cp_size = cp_group.size() rolled_tensor = tensor.clone() @@ -1302,6 +1321,7 @@ def __init__( pg_collection=pg_collection, is_mtp_layer=True, mtp_layer_number=self.layer_number, + cp_stage_entry_partition_mode="zigzag", name=(name + ".mtp_model_layer") if name is not None else None, ) elif self.config.mtp_num_layers is not None: @@ -1818,8 +1838,18 @@ def forward( [s, b, h], and optionally the updated context tensor if cross-attention is used. """ assert context is None, "multi token prediction + cross attention is not yet supported." - _orig_cp_group = self.cp_group - self.cp_group = resolve_cp_group(self.cp_group, packed_seq_params) + cp_partition_mode = getattr(packed_seq_params, 'cp_partition_mode', None) + if cp_partition_mode is not None: + preferred_partition_mode = get_preferred_cp_partition_mode_for_layer( + self.mtp_model_layer, self.config + ) + if preferred_partition_mode is not None and preferred_partition_mode != cp_partition_mode: + raise NotImplementedError( + "MTP inner layer CP partition mode preference does not match the current " + "MTP input layout: " + f"preferred {preferred_partition_mode!r}, got {cp_partition_mode!r}." + ) + input_ids, position_ids, padding_mask, decoder_input, hidden_states = self._get_embeddings( input_ids=input_ids, position_ids=position_ids, @@ -1863,8 +1893,6 @@ def forward( packed_seq_params=packed_seq_params, sequence_len_offset=sequence_len_offset, ) - - self.cp_group = _orig_cp_group return hidden_states, input_ids, position_ids, padding_mask def sharded_state_dict( @@ -1884,6 +1912,8 @@ def sharded_state_dict( """ sharded_state_dict = super().sharded_state_dict(prefix, sharded_offsets, metadata) + # TODO(yuzhongw): audit remaining Mamba-only wording in + # MTP comments once legacy ``mamba_submodules`` compatibility is removed. # Backward compatibility: GPT MTP checkpoints were saved with the submodule # named 'transformer_layer'. Remap checkpoint keys so old checkpoints load # correctly. Mamba MTP models keep 'mtp_model_layer' as their native format diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index 142fba9a69f..1e1e3416bd4 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -10,6 +10,13 @@ from torch import Tensor from megatron.core import parallel_state, tensor_parallel +from megatron.core.context_parallel_layout import ( + CpPartitionMode, + CpPartitionModeConverter, + get_preferred_cp_partition_mode_for_layer, + get_stage_entry_partition_mode, + replace_packed_seq_params_cp_partition_mode, +) from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import replace_prefix_for_sharding from megatron.core.enums import Fp8Recipe @@ -19,7 +26,7 @@ from megatron.core.fusions.fused_layer_norm import FusedLayerNorm from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.inference.utils import InferenceMode -from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group from megatron.core.pipeline_parallel.utils import is_vp_first_stage, is_vp_last_stage from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import CheckpointManager @@ -288,6 +295,7 @@ def __init__( post_process: bool = True, pg_collection: Optional[ProcessGroupCollection] = None, vp_stage: Optional[int] = None, + cp_stage_entry_partition_mode: Optional[str] = None, ): super().__init__(config=config) @@ -304,6 +312,7 @@ def __init__( self.pre_process = pre_process self.post_process = post_process self.vp_stage = vp_stage + self.cp_stage_entry_partition_mode = cp_stage_entry_partition_mode # required for pipeline parallel schedules self.input_tensor = None @@ -578,6 +587,65 @@ def _setup_fused_tp_communication(self): def _get_layer(self, layer_number: int): return self.layers[layer_number] + def _convert_cp_partition_mode_for_layer( + self, + *, + local_index: int, + current_partition_mode: CpPartitionMode, + hidden_states: Tensor, + attention_mask: Optional[Tensor], + attention_bias: Optional[Tensor], + packed_seq_params: Optional[PackedSeqParams], + padding_mask: Optional[Tensor], + input_ids: Optional[Tensor], + preferred_partition_mode: Optional[CpPartitionMode], + ): + """Convert per-token tensors to the layout preferred by one local layer.""" + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + if cp_group is None or cp_group.size() <= 1: + return hidden_states, padding_mask, input_ids + if preferred_partition_mode is None or preferred_partition_mode == current_partition_mode: + return hidden_states, padding_mask, input_ids + + current_to_preferred_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode=current_partition_mode, + target_partition_mode=preferred_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + if not current_to_preferred_converter.conversion_needed: + return hidden_states, padding_mask, input_ids + + current_to_preferred_converter.assert_no_dense_attention_inputs( + attention_mask=attention_mask, + attention_bias=attention_bias, + hidden_states=hidden_states, + ) + hidden_states = current_to_preferred_converter.convert( + hidden_states, + seq_dim=0, + sequence_parallel=self.config.sequence_parallel, + ) + # Model-level rotary_pos_emb is only consumed by regular attention, whose + # CP layout preference is zigzag. MLA/CSA/DSv4-style variants must ignore + # external RoPE and manage any RoPE positions internally, so this layout + # edge intentionally does not convert rotary_pos_emb. + if padding_mask is not None: + padding_mask = current_to_preferred_converter.convert( + padding_mask, + seq_dim=1, + sequence_parallel=self.config.sequence_parallel, + ) + if input_ids is not None: + input_ids = current_to_preferred_converter.convert( + input_ids, + seq_dim=1, + ) + + return hidden_states, padding_mask, input_ids + def _checkpointed_forward( self, hidden_states: Tensor, @@ -610,6 +678,22 @@ def _checkpointed_forward( if extract_layer_indices is None: extract_layer_indices = set() intermediate_hidden_states: List[Tensor] = [] + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + cp_layout_needed = ( + cp_group is not None + and cp_group.size() > 1 + and self.config.cp_partition_mode == "auto" + ) + stage_entry_partition_mode = ( + get_stage_entry_partition_mode( + packed_seq_params, + self.cp_stage_entry_partition_mode, + owner_name=type(self).__name__, + cp_group=cp_group, + ) + if cp_layout_needed + else None + ) def custom(start: int, end: int): def custom_forward( @@ -620,8 +704,77 @@ def custom_forward( rotary_pos_emb, padding_mask=None, ): + current_partition_mode = stage_entry_partition_mode + if cp_layout_needed: + for index in range(start): + layer = self._get_layer(index) + preferred_partition_mode = get_preferred_cp_partition_mode_for_layer( + layer, getattr(layer, "config", self.config) + ) + if preferred_partition_mode is not None: + current_partition_mode = preferred_partition_mode + local_packed_seq_params = packed_seq_params + if cp_layout_needed: + local_packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, current_partition_mode + ) + local_input_ids = input_ids + if cp_layout_needed and current_partition_mode != stage_entry_partition_mode: + chunk_entry_converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=local_packed_seq_params, + source_partition_mode=stage_entry_partition_mode, + target_partition_mode=current_partition_mode, + config=self.config, + tp_group=self.pg_collection.tp, + ) + chunk_entry_converter.assert_no_dense_attention_inputs( + attention_mask=attention_mask, + attention_bias=attention_bias, + hidden_states=hidden_states, + ) + # Checkpoint chunks may enter after earlier layers changed the + # current CP layout. Re-align batch/token side tensors here, but + # keep model-level RoPE in the regular-attention zigzag layout. + if padding_mask is not None: + padding_mask = chunk_entry_converter.convert( + padding_mask, + seq_dim=1, + sequence_parallel=self.config.sequence_parallel, + ) + if local_input_ids is not None: + local_input_ids = chunk_entry_converter.convert( + local_input_ids, + seq_dim=1, + ) for index in range(start, end): layer = self._get_layer(index) + if cp_layout_needed: + preferred_partition_mode = get_preferred_cp_partition_mode_for_layer( + layer, getattr(layer, "config", self.config) + ) + (hidden_states, padding_mask, local_input_ids) = ( + self._convert_cp_partition_mode_for_layer( + local_index=index, + current_partition_mode=current_partition_mode, + hidden_states=hidden_states, + attention_mask=attention_mask, + attention_bias=attention_bias, + packed_seq_params=local_packed_seq_params, + padding_mask=padding_mask, + input_ids=local_input_ids, + preferred_partition_mode=preferred_partition_mode, + ) + ) + if preferred_partition_mode is not None: + local_packed_seq_params = replace_packed_seq_params_cp_partition_mode( + local_packed_seq_params, preferred_partition_mode + ) + current_partition_mode = getattr( + local_packed_seq_params, + "cp_partition_mode", + preferred_partition_mode or current_partition_mode, + ) # Get appropriate inner quantization context if use_inner_quantization_context: @@ -648,9 +801,9 @@ def custom_forward( rotary_pos_emb=rotary_pos_emb, attention_bias=attention_bias, inference_context=None, - packed_seq_params=packed_seq_params, + packed_seq_params=local_packed_seq_params, padding_mask=padding_mask, - input_ids=input_ids, + input_ids=local_input_ids, ) return hidden_states, context @@ -950,6 +1103,23 @@ def forward( ) with rng_context, outer_quantization_context: + cp_group = resolve_cp_group(self.pg_collection.cp, packed_seq_params) + cp_layout_needed = ( + cp_group is not None + and cp_group.size() > 1 + and self.config.cp_partition_mode == "auto" + ) + current_partition_mode = None + if cp_layout_needed: + current_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + self.cp_stage_entry_partition_mode, + owner_name=type(self).__name__, + cp_group=cp_group, + ) + packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, current_partition_mode + ) # Forward pass. if self.config.recompute_granularity == 'full' and self.training: checkpointed_result = self._checkpointed_forward( @@ -975,6 +1145,33 @@ def forward( hidden_states = checkpointed_result else: for l_no, layer in enumerate(self.layers): + if cp_layout_needed: + preferred_partition_mode = get_preferred_cp_partition_mode_for_layer( + layer, getattr(layer, "config", self.config) + ) + (hidden_states, padding_mask, input_ids) = ( + self._convert_cp_partition_mode_for_layer( + local_index=l_no, + current_partition_mode=current_partition_mode, + hidden_states=hidden_states, + attention_mask=attention_mask, + attention_bias=attention_bias, + packed_seq_params=packed_seq_params, + padding_mask=padding_mask, + input_ids=input_ids, + preferred_partition_mode=preferred_partition_mode, + ) + ) + if preferred_partition_mode is not None: + packed_seq_params = replace_packed_seq_params_cp_partition_mode( + packed_seq_params, preferred_partition_mode + ) + current_partition_mode = getattr( + packed_seq_params, + "cp_partition_mode", + preferred_partition_mode or current_partition_mode, + ) + # Get appropriate inner quantization context if use_inner_quantization_context: if self.config.fp8: diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 4091b647103..6dd388542b9 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -309,7 +309,7 @@ class TransformerConfig(ModelParallelConfig): ) """Type of attention variant to use. Currently support gated_delta_net, dsa, and dsv4_hybrid.""" - cp_partition_mode: Literal["zigzag", "contiguous"] = "zigzag" + cp_partition_mode: Literal["zigzag", "contiguous", "auto"] = "zigzag" """How THD sequence rows are partitioned across context-parallel ranks.""" experimental_attention_variant_loss_scale_func: Optional[Callable[[torch.Tensor], None]] = None @@ -1553,22 +1553,48 @@ def __post_init__(self): self.experimental_attention_variant = self.linear_attention_type self.linear_attention_type = None - if self.cp_partition_mode not in ("zigzag", "contiguous"): + if self.cp_partition_mode not in ("zigzag", "contiguous", "auto"): raise ValueError(f"Unsupported cp_partition_mode: {self.cp_partition_mode}") - + if self.cp_partition_mode == "auto" and self.overlap_moe_expert_parallel_comm: + raise ValueError( + 'cp_partition_mode="auto" is not supported with ' + "overlap_moe_expert_parallel_comm in this rollout." + ) if self.context_parallel_size > 1: - if ( - self.experimental_attention_variant == "dsv4_hybrid" - and self.cp_partition_mode != "contiguous" - ): - raise ValueError("DSv4 Hybrid with CP requires cp_partition_mode='contiguous'.") - if ( - self.experimental_attention_variant != "dsv4_hybrid" - and self.cp_partition_mode != "zigzag" - ): - raise ValueError( - "cp_partition_mode='contiguous' currently is only supported with dsv4_hybrid." - ) + if self.cp_partition_mode == "contiguous": + if ( + self.multi_latent_attention + and self.experimental_attention_variant != "dsv4_hybrid" + ): + raise ValueError( + "cp_partition_mode='contiguous' is not supported with " + "multi_latent_attention outside dsv4_hybrid." + ) + if self.experimental_attention_variant not in ( + "dsv4_hybrid", + "gated_delta_net", + ): + raise ValueError( + "cp_partition_mode='contiguous' with context parallelism currently " + "requires experimental_attention_variant to be either 'dsv4_hybrid' " + "or 'gated_delta_net'." + ) + if ( + self.experimental_attention_variant == "gated_delta_net" + and self.linear_cp_mode == "headwise" + ): + raise ValueError( + "cp_partition_mode='contiguous' is incompatible with " + "gated_delta_net linear_cp_mode='headwise'." + ) + elif self.cp_partition_mode == "zigzag": + if self.experimental_attention_variant == "dsv4_hybrid": + raise ValueError( + "DSv4 Hybrid with context parallelism requires " + "cp_partition_mode='contiguous'." + ) + elif self.cp_partition_mode != "auto": + raise ValueError(f"Unsupported cp_partition_mode: {self.cp_partition_mode}") # Normalize the deprecated DSv4 kernel switch only after all deprecated attention # selectors have been folded into experimental_attention_variant, and immediately diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index ea714552464..df8dd658804 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -16,6 +16,7 @@ from torch import Tensor from megatron.core import parallel_state, tensor_parallel +from megatron.core.context_parallel_layout import get_preferred_cp_partition_mode_for_layer from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import apply_prefix_mapping from megatron.core.inference.utils import InferenceMode @@ -1373,13 +1374,23 @@ def _reconstruct_packed_seq_params_from_kwargs(self, kwargs): the padded static value and cannot be read from CUDA tensors during graph capture. pad_between_seqs is conservatively enabled because its Python value is not a graph input and inferring it from CUDA tensors would synchronize during capture. + cp_partition_mode is inferred from this layer's layout preference instead + of the deprecated TransformerConfig.cp_partition_mode global field. """ if 'cu_seqlens_q' not in kwargs: return max_seqlen = self.config.max_seqlen_per_dp_cp_rank * self.config.context_parallel_size + cp_partition_mode = get_preferred_cp_partition_mode_for_layer(self, self.config) + if cp_partition_mode is None: + if self.config.context_parallel_size > 1 or self.config.dynamic_context_parallel: + raise ValueError( + "Cannot reconstruct THD PackedSeqParams for a layout-agnostic transformer " + "layer under context parallelism. The CP partition mode must be provided " + "by the model-level layout plan." + ) packed_seq_params = PackedSeqParams( qkv_format='thd', - cp_partition_mode=self.config.cp_partition_mode, + cp_partition_mode=cp_partition_mode, cu_seqlens_q=kwargs.pop('cu_seqlens_q'), cu_seqlens_kv=kwargs.pop('cu_seqlens_kv'), cu_seqlens_q_padded=kwargs.pop('cu_seqlens_q_padded'), diff --git a/megatron/core/utils.py b/megatron/core/utils.py index ac97916d21b..40aa7520800 100644 --- a/megatron/core/utils.py +++ b/megatron/core/utils.py @@ -42,6 +42,10 @@ HAVE_DTENSOR = False from megatron.core import parallel_state +from megatron.core.context_parallel_layout import ( + get_context_parallel_layout_chunk_indices, + get_thd_context_parallel_rank_indices, +) from megatron.core.dist_checkpointing.mapping import ShardedTensor from megatron.core.packed_seq_params import PackedSeqParams @@ -2337,7 +2341,9 @@ def _broadcast_cu_seqlens(): def _get_batch_on_this_cp_rank_per_document_balancing( - batch: dict[str, torch.Tensor], cp_group: torch.distributed.ProcessGroup + batch: dict[str, torch.Tensor], + cp_group: torch.distributed.ProcessGroup, + cp_partition_mode: Optional[str] = "zigzag", ): """Partition a batch across CP ranks with per-document zigzag load balancing. @@ -2363,22 +2369,32 @@ def _get_batch_on_this_cp_rank_per_document_balancing( cp_rank = torch.distributed.get_rank(cp_group) if cp_size > 1: - # cu_seqlens / cu_seqlens_padded carry a leading batch dim (1, n). - # tex.thd_get_partitioned_indices expects a 1-D tensor, so squeeze - # the batch dim inline without mutating the batch dict. + if cp_partition_mode is None: + raise ValueError( + "cp_partition_mode must be provided when partitioning an SFT batch " + "across context-parallel ranks." + ) + # cu_seqlens / cu_seqlens_padded carry the dataloader's batch dim (1, n). + # tex.thd_get_partitioned_indices expects a 1-D tensor, so squeeze the + # batch dim inline without mutating the batch dict. cu_seqlens_for_te = ( batch["cu_seqlens_padded"] if batch["cu_seqlens_padded"] is not None else batch["cu_seqlens"] )[0] - index = tex.thd_get_partitioned_indices( - cu_seqlens_for_te, - ( - batch["tokens"].size(1) if batch["tokens"] is not None else batch["labels"].size(1) - ), # NOTE(asolergi-nv): Labels to enable PP! - cp_size, - cp_rank, - ) + total_tokens = ( + batch["tokens"].size(1) if batch["tokens"] is not None else batch["labels"].size(1) + ) # NOTE(asolergi-nv): Labels to enable PP! + if cp_partition_mode == "zigzag": + index = tex.thd_get_partitioned_indices( + cu_seqlens_for_te, total_tokens, cp_size, cp_rank + ) + elif cp_partition_mode == "contiguous": + index = get_thd_context_parallel_rank_indices( + cu_seqlens_for_te, cp_size, cp_rank, cp_partition_mode + ) + else: + raise ValueError(f"Unsupported context-parallel partition mode {cp_partition_mode!r}.") SEQUENCE_KEYS = ('tokens', 'labels', 'loss_mask', 'position_ids') for key in SEQUENCE_KEYS: if batch.get(key) is not None: @@ -2387,7 +2403,9 @@ def _get_batch_on_this_cp_rank_per_document_balancing( def _get_batch_on_this_cp_rank_per_sequence_balancing( - batch: dict[str, torch.Tensor], cp_group: torch.distributed.ProcessGroup + batch: dict[str, torch.Tensor], + cp_group: torch.distributed.ProcessGroup, + cp_partition_mode: Optional[str] = "zigzag", ): """Partition a batch across CP ranks with per-sequence zigzag load balancing. @@ -2412,7 +2430,6 @@ def _get_batch_on_this_cp_rank_per_sequence_balancing( dict[str, torch.Tensor]: The batch with sequence-dimension tensors partitioned to this CP rank. """ - cp_size = torch.distributed.get_world_size(cp_group) cp_rank = torch.distributed.get_rank(cp_group) @@ -2427,6 +2444,11 @@ def _get_batch_on_this_cp_rank_per_sequence_balancing( ) if cp_size > 1: + if cp_partition_mode is None: + raise ValueError( + "cp_partition_mode must be provided when partitioning a pretraining batch " + "across context-parallel ranks." + ) for key, val in batch.items(): if key in METADATA_KEYS or val is None: continue @@ -2437,9 +2459,9 @@ def _get_batch_on_this_cp_rank_per_sequence_balancing( val.shape[seq_dim] // (2 * cp_size), *val.shape[(seq_dim + 1) :], ) - index = torch.zeros(2, dtype=torch.int64, device=val.device) - index[0].fill_(cp_rank) - index[1].fill_(2 * cp_size - cp_rank - 1) + index = get_context_parallel_layout_chunk_indices( + cp_size, cp_rank, cp_partition_mode + ).to(device=val.device) val = val.index_select(seq_dim, index) val = val.view(*val.shape[0:seq_dim], -1, *val.shape[(seq_dim + 2) :]) batch[key] = val @@ -2548,6 +2570,7 @@ def get_batch_on_this_cp_rank( cp_group: Optional[torch.distributed.ProcessGroup] = None, hybrid_cp_group_func: Optional[Callable[[int], torch.distributed.ProcessGroup]] = None, use_per_sequence_balancing: bool = False, + cp_partition_mode: Optional[str] = "zigzag", ): """Dispatch batch partitioning across context-parallel ranks. @@ -2576,6 +2599,8 @@ def get_batch_on_this_cp_rank( even when ``cu_seqlens`` is present (e.g., for inter-document masking where document lengths are not divisible by ``2 * cp_size``). + cp_partition_mode (Optional[str]): Runtime CP layout used to partition + sequence tensors. Defaults to ``"zigzag"`` for legacy callers. Returns: Dict[str, Any]: The batch with sequence-dimension tensors partitioned @@ -2589,7 +2614,9 @@ def get_batch_on_this_cp_rank( cp_group = parallel_state.get_context_parallel_group() if use_per_sequence_balancing or batch.get("cu_seqlens") is None: - batch = _get_batch_on_this_cp_rank_per_sequence_balancing(batch, cp_group=cp_group) + batch = _get_batch_on_this_cp_rank_per_sequence_balancing( + batch, cp_group=cp_group, cp_partition_mode=cp_partition_mode + ) elif is_hybrid_cp: assert ( batch['local_cp_size'] is not None @@ -2597,11 +2624,13 @@ def get_batch_on_this_cp_rank( if batch['local_cp_size'].item() > 1: hybrid_cp_group = hybrid_cp_group_func(group_size=batch['local_cp_size'].item()) batch = _get_batch_on_this_cp_rank_per_sequence_balancing( - batch, cp_group=hybrid_cp_group + batch, cp_group=hybrid_cp_group, cp_partition_mode=cp_partition_mode ) batch["hybrid_cp_group"] = hybrid_cp_group else: - batch = _get_batch_on_this_cp_rank_per_document_balancing(batch, cp_group=cp_group) + batch = _get_batch_on_this_cp_rank_per_document_balancing( + batch, cp_group=cp_group, cp_partition_mode=cp_partition_mode + ) return batch @@ -2612,6 +2641,7 @@ def get_thd_batch_on_this_cp_rank( max_seqlen: torch.Tensor, cp_size: Optional[int] = None, cp_rank: Optional[int] = None, + cp_partition_mode: Optional[str] = "zigzag", ): """Slice each sub-sample in a packed sample batch input along sequence dimension into multiple chunks, which are parallelized @@ -2625,18 +2655,31 @@ def get_thd_batch_on_this_cp_rank( cu_seqlens_kv_padded=cu_seqlens_padded, max_seqlen_q=int(max_seqlen[0].item()), max_seqlen_kv=int(max_seqlen[0].item()), + cp_partition_mode=cp_partition_mode, ) cp_size = parallel_state.get_context_parallel_world_size() if cp_size is None else cp_size cp_rank = parallel_state.get_context_parallel_rank() if cp_rank is None else cp_rank if cp_size > 1: # slice batch along sequence dimension for context parallelism + if cp_partition_mode is None: + raise ValueError( + "cp_partition_mode must be provided when partitioning a THD batch " + "across context-parallel ranks." + ) assert tex is not None and is_te_min_version("1.10.0"), ( "Please update Transformer Engine to >= 1.10 to use " "Context Parallel with THD format data" ) - index = tex.thd_get_partitioned_indices( - cu_seqlens_padded, batch['tokens'].size(1), cp_size, cp_rank - ) + if cp_partition_mode == "zigzag": + index = tex.thd_get_partitioned_indices( + cu_seqlens_padded, batch['tokens'].size(1), cp_size, cp_rank + ) + elif cp_partition_mode == "contiguous": + index = get_thd_context_parallel_rank_indices( + cu_seqlens_padded, cp_size, cp_rank, cp_partition_mode + ) + else: + raise ValueError(f"Unsupported context-parallel partition mode {cp_partition_mode!r}.") for key, data in batch.items(): if key in {'attention_mask', 'cu_seqlens', 'cu_seqlens_padded', 'max_seqlen'}: continue diff --git a/megatron/elastification/pretrain_hybrid_flex.py b/megatron/elastification/pretrain_hybrid_flex.py index be22082127e..51ada2e5cea 100644 --- a/megatron/elastification/pretrain_hybrid_flex.py +++ b/megatron/elastification/pretrain_hybrid_flex.py @@ -26,7 +26,6 @@ get_dynamic_data_context_parallel_groups, get_pipeline_model_parallel_rank, get_pipeline_model_parallel_world_size, - get_tensor_model_parallel_group, get_tensor_model_parallel_rank, ) from megatron.core.rerun_state_machine import get_rerun_state_machine @@ -115,6 +114,12 @@ def model_provider( HybridModel: The returned model """ args = get_args() + if getattr(args, "cp_partition_mode", "zigzag") == "auto": + raise ValueError( + 'megatron/elastification/pretrain_hybrid_flex.py does not support ' + 'cp_partition_mode="auto"; use pretrain_hybrid.py for the CP layout ' + 'refactor path.' + ) if has_nvidia_modelopt: model = model_provider_modelopt( diff --git a/megatron/training/models/gpt.py b/megatron/training/models/gpt.py index e418e9f7368..7cd2d06de92 100644 --- a/megatron/training/models/gpt.py +++ b/megatron/training/models/gpt.py @@ -9,6 +9,7 @@ from megatron.core.distributed.distributed_data_parallel_config import DistributedDataParallelConfig from megatron.core.enums import ModelType from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_stage_input_cp_partition_mode, get_transformer_block_with_experimental_attention_variant_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel @@ -23,6 +24,7 @@ ) from megatron.core.post_training.modelopt.gpt.model_specs import get_gpt_modelopt_spec from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.utils import get_pg_rank from megatron.core.transformer.dot_product_attention import ( DotProductAttention as MCoreDotProductAttention, ) @@ -334,6 +336,20 @@ def build_model( vp_stage=vp_stage, vp_size=vp_size ) and is_pp_last_stage(pg_collection.pp) + cp_stage_entry_partition_mode = ( + "zigzag" + if self._model_config.transformer.cp_partition_mode == "auto" + else self._model_config.transformer.cp_partition_mode + ) + if self._model_config.transformer.cp_partition_mode == "auto": + cp_stage_entry_partition_mode = ( + get_experimental_attention_variant_stage_input_cp_partition_mode( + self._model_config.transformer, + vp_stage=vp_stage, + pp_rank=get_pg_rank(pg_collection.pp), + ) + ) + model = GPTModel( config=self._model_config.transformer, transformer_layer_spec=transformer_layer_spec, @@ -354,6 +370,7 @@ def build_model( post_process=post_process, pg_collection=pg_collection, vp_stage=vp_stage, + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, ) return model diff --git a/megatron/training/models/hybrid.py b/megatron/training/models/hybrid.py index 397b762a4dc..21c827aaa05 100644 --- a/megatron/training/models/hybrid.py +++ b/megatron/training/models/hybrid.py @@ -6,6 +6,9 @@ from megatron.core.distributed.distributed_data_parallel_config import DistributedDataParallelConfig from megatron.core.enums import ModelType +from megatron.core.models.hybrid.hybrid_layer_allocation import ( + get_hybrid_stage_input_cp_partition_mode_for_stage, +) from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_inference_stack_spec from megatron.core.models.hybrid.hybrid_layer_specs import ( hybrid_stack_spec as default_hybrid_stack_spec, @@ -175,6 +178,22 @@ def build_model( post_process = ( post_process if post_process is not None else is_pp_last_stage(pg_collection.pp) ) + cp_stage_entry_partition_mode = self._model_config.transformer.cp_partition_mode + if self._model_config.transformer.cp_partition_mode == "auto": + cp_stage_entry_partition_mode = ( + get_hybrid_stage_input_cp_partition_mode_for_stage( + self._model_config.transformer, + self._model_config.hybrid_layer_pattern, + pg_collection.pp, + vp_stage, + first_stage_layers=( + self._model_config.transformer.num_layers_in_first_pipeline_stage + ), + last_stage_layers=( + self._model_config.transformer.num_layers_in_last_pipeline_stage + ), + ) + ) return HybridModel( config=self._model_config.transformer, hybrid_stack_spec=hybrid_stack_spec, @@ -192,6 +211,7 @@ def build_model( post_process=post_process, pg_collection=pg_collection, vp_stage=vp_stage, + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, ) def build_distributed_models( diff --git a/megatron/training/training.py b/megatron/training/training.py index 4ad7f00c4db..49c3c65e49b 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -2718,6 +2718,7 @@ def dummy_train_step(data_iterator): is_hybrid_cp=is_hybrid_cp, cp_group=get_context_parallel_group(), hybrid_cp_group_func=get_dynamic_data_context_parallel_groups, + cp_partition_mode=getattr(args, 'cp_partition_mode', None), ) diff --git a/pretrain_gpt.py b/pretrain_gpt.py index 9e99200b27a..64b565f0fd2 100644 --- a/pretrain_gpt.py +++ b/pretrain_gpt.py @@ -25,6 +25,10 @@ from gpt_builders import gpt_builder from megatron.core import mpu +from megatron.core.context_parallel_layout import ( + get_stage_entry_partition_mode, + prebuild_thd_cp_partition_routes, +) from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder from megatron.core.datasets.data_schedule import get_batch_on_this_rank_for_sequence_packing from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset @@ -34,6 +38,7 @@ PackedSeqParams, get_thd_padding_kwargs, pad_sequence_for_thd, + resolve_cp_group, resolve_thd_tail_padding_policy, ) from megatron.core.rerun_state_machine import get_rerun_state_machine @@ -76,7 +81,11 @@ stimer = StragglerDetector() -def get_batch(data_iterator, vp_stage: Optional[int] = None): +def get_batch( + data_iterator, + vp_stage: Optional[int] = None, + cp_partition_mode="zigzag", +): """Generate a batch. Packed sequence support (SFT / ``--sft`` flag): @@ -105,6 +114,9 @@ def get_batch(data_iterator, vp_stage: Optional[int] = None): regardless of pipeline stage. Difference from ``pretrain_hybrid.py``: + - TODO(yuzhongw): this comparison still uses historical + Mamba wording. Re-audit and rename to Hybrid where the behavior is + no longer Mamba-specific outside this PR. - Return format: GPT returns a 6-tuple ``(tokens, labels, loss_mask, attention_mask, position_ids, packed_seq_params)`` where ``packed_seq_params`` is a @@ -132,14 +144,32 @@ def get_batch(data_iterator, vp_stage: Optional[int] = None): if args.sequence_packing_scheduler is not None: # `get_batch_on_this_rank_for_sequence_packing` owns scheduler THD metadata # and returns a 7-tuple including `padding_mask`. - return get_batch_on_this_rank_for_sequence_packing( + batch = get_batch_on_this_rank_for_sequence_packing( data_iterator, vpp_size=config.virtual_pipeline_model_parallel_size, mtp_on_this_rank=mtp_on_this_rank(config, ignore_virtual=False, vp_stage=vp_stage), vp_stage=vp_stage, dynamic_cp=args.dynamic_context_parallel, config=config, + cp_partition_mode=cp_partition_mode, ) + packed_seq_params = batch[5] + if packed_seq_params is not None: + packed_seq_params.cp_group = resolve_cp_group( + mpu.get_context_parallel_group(), packed_seq_params + ) + packed_seq_params.cp_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + cp_partition_mode, + owner_name="pretrain_gpt.get_batch", + cp_group=packed_seq_params.cp_group, + ) + if getattr(config, "experimental_attention_variant", None) not in (None, "none"): + prebuild_thd_cp_partition_routes( + packed_seq_params, + resolve_cp_group(mpu.get_context_parallel_group(), packed_seq_params), + ) + return batch # TODO: this is pretty hacky, find a better way is_packed_sequence = args.sft or (args.use_varlen_dataset and not args.varlen_sbhd_validation) @@ -174,30 +204,53 @@ def get_batch(data_iterator, vp_stage: Optional[int] = None): # For middle pipeline stages with packed sequences, only cu_seqlens and # max_seqlen are needed (for attention masking); skip the full batch. if not is_first_or_last_pipeline_stage(vp_stage) and is_packed_sequence: + packed_seq_params = PackedSeqParams( + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + max_seqlen_q=int(max_seqlen[0].item()), + max_seqlen_kv=int(max_seqlen[0].item()), + qkv_format='thd', + cp_partition_mode=cp_partition_mode, + ) + if packed_seq_params is not None: + packed_seq_params.cp_group = resolve_cp_group( + mpu.get_context_parallel_group(), packed_seq_params + ) + packed_seq_params.cp_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + cp_partition_mode, + owner_name="pretrain_gpt.get_batch", + cp_group=packed_seq_params.cp_group, + ) + if getattr(config, "experimental_attention_variant", None) not in (None, "none"): + prebuild_thd_cp_partition_routes( + packed_seq_params, + resolve_cp_group(mpu.get_context_parallel_group(), packed_seq_params), + ) return ( None, None, None, None, None, - PackedSeqParams( - cu_seqlens_q=cu_seqlens, - cu_seqlens_kv=cu_seqlens, - max_seqlen_q=int(max_seqlen[0].item()), - max_seqlen_kv=int(max_seqlen[0].item()), - qkv_format='thd', - ), + packed_seq_params, None, ) thd_tail_padding_policy = resolve_thd_tail_padding_policy(config) if cu_seqlens is None: # slice batch along sequence dimension for context parallelism - batch = get_batch_on_this_cp_rank(batch) # The implementation of this function is in MCore + batch = get_batch_on_this_cp_rank( # The implementation of this function is in MCore + batch, cp_partition_mode=cp_partition_mode + ) packed_seq_params = None else: # Packed THD format batch, packed_seq_params = get_thd_batch_on_this_cp_rank( - batch, cu_seqlens, cu_seqlens_padded, max_seqlen + batch, + cu_seqlens, + cu_seqlens_padded, + max_seqlen, + cp_partition_mode=cp_partition_mode, ) # Pad the already-packed THD tensors at the end when requested. A configured @@ -239,6 +292,22 @@ def get_batch(data_iterator, vp_stage: Optional[int] = None): if 'position_ids' in batch: batch['position_ids'] = position_ids + if packed_seq_params is not None: + packed_seq_params.cp_group = resolve_cp_group( + mpu.get_context_parallel_group(), packed_seq_params + ) + packed_seq_params.cp_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + cp_partition_mode, + owner_name="pretrain_gpt.get_batch", + cp_group=packed_seq_params.cp_group, + ) + if getattr(config, "experimental_attention_variant", None) not in (None, "none"): + prebuild_thd_cp_partition_routes( + packed_seq_params, + resolve_cp_group(mpu.get_context_parallel_group(), packed_seq_params), + ) + # Unpack explicitly to avoid relying on dict insertion order. return ( batch.get('tokens'), @@ -365,8 +434,10 @@ def forward_step(data_iterator, model: GPTModel, return_schedule_plan: bool = Fa global stimer with stimer(bdata=True): vp_stage = get_attr_wrapped_model(model, "vp_stage") + decoder = get_attr_wrapped_model(model, "decoder") + cp_partition_mode = decoder.cp_stage_entry_partition_mode tokens, labels, loss_mask, attention_mask, position_ids, packed_seq_params, padding_mask = ( - get_batch(data_iterator, vp_stage) + get_batch(data_iterator, vp_stage, cp_partition_mode=cp_partition_mode) ) timers('batch-generator').stop() diff --git a/pretrain_hybrid.py b/pretrain_hybrid.py index 5fdc7d8645a..d79cde4de95 100644 --- a/pretrain_hybrid.py +++ b/pretrain_hybrid.py @@ -24,12 +24,16 @@ from hybrid_builders import hybrid_builder from megatron.core import mpu +from megatron.core.context_parallel_layout import ( + get_stage_entry_partition_mode, + prebuild_thd_cp_partition_routes, +) from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder from megatron.core.datasets.data_schedule import get_batch_on_this_rank_for_sequence_packing from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset from megatron.core.enums import ModelType from megatron.core.models.hybrid.hybrid_model import HybridModel -from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group from megatron.core.parallel_state import ( get_context_parallel_group, get_dynamic_data_context_parallel_groups, @@ -76,7 +80,56 @@ stimer = StragglerDetector() -def get_batch(data_iterator, vp_stage=None): +def _build_packed_seq_params_for_batch(batch, args, cp_partition_mode): + cu_seqlens = batch.get('cu_seqlens', None) + if cu_seqlens is None: + return None + + if cu_seqlens.dim() == 2: + cu_seqlens = cu_seqlens.squeeze(0) + cu_seqlens_padded = batch.get('cu_seqlens_padded', None) + if cu_seqlens_padded is not None and cu_seqlens_padded.dim() == 2: + cu_seqlens_padded = cu_seqlens_padded.squeeze(0) + max_seqlen = batch.get('max_seqlen') + if max_seqlen is not None and max_seqlen.dim() > 0: + max_seqlen = max_seqlen.squeeze(0) + + cu_seqlens_for_params = cu_seqlens_padded if cu_seqlens_padded is not None else cu_seqlens + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu_seqlens_for_params, + cu_seqlens_kv=cu_seqlens_for_params, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=int(max_seqlen.item()), + max_seqlen_kv=int(max_seqlen.item()), + local_cp_size=( + int(batch['local_cp_size'].item()) + if batch.get('local_cp_size', None) is not None + else None + ), + cp_group=batch.get('hybrid_cp_group', None), + total_tokens=int(cu_seqlens_for_params[-1].item()), + tokens_per_sample=args.seq_length, + cp_partition_mode=cp_partition_mode, + ) + packed_seq_params.cp_group = resolve_cp_group( + get_context_parallel_group(), packed_seq_params + ) + packed_seq_params.cp_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + cp_partition_mode, + owner_name="pretrain_hybrid.get_batch", + cp_group=packed_seq_params.cp_group, + ) + prebuild_thd_cp_partition_routes( + packed_seq_params, + resolve_cp_group(get_context_parallel_group(), packed_seq_params), + ) + return packed_seq_params + + +def get_batch(data_iterator, vp_stage=None, cp_partition_mode="zigzag"): """Generate a batch.""" batch_keys = [ @@ -124,6 +177,21 @@ def get_batch(data_iterator, vp_stage=None): vp_stage=vp_stage, dynamic_cp=is_dynamic_cp, config=config, + cp_partition_mode=cp_partition_mode, + ) + if packed_seq_params is not None: + packed_seq_params.cp_group = resolve_cp_group( + get_context_parallel_group(), packed_seq_params + ) + packed_seq_params.cp_partition_mode = get_stage_entry_partition_mode( + packed_seq_params, + cp_partition_mode, + owner_name="pretrain_hybrid.get_batch", + cp_group=packed_seq_params.cp_group, + ) + prebuild_thd_cp_partition_routes( + packed_seq_params, + resolve_cp_group(get_context_parallel_group(), packed_seq_params), ) return ( attention_mask, @@ -178,6 +246,7 @@ def get_batch(data_iterator, vp_stage=None): if not is_first_or_last_pipeline_stage(vp_stage) and not mtp_on_this_rank: assert has_cu_seqlens + packed_seq_params = _build_packed_seq_params_for_batch(batch, args, cp_partition_mode) return ( None, batch['cu_seqlens'], @@ -190,7 +259,7 @@ def get_batch(data_iterator, vp_stage=None): None, None, None, - None, + packed_seq_params, ) batch = get_batch_on_this_cp_rank( @@ -199,11 +268,13 @@ def get_batch(data_iterator, vp_stage=None): cp_group=get_context_parallel_group(), hybrid_cp_group_func=get_dynamic_data_context_parallel_groups, use_per_sequence_balancing=args.dataloader_inter_document_masking and not is_sft, + cp_partition_mode=cp_partition_mode, ) + packed_seq_params = _build_packed_seq_params_for_batch(batch, args, cp_partition_mode) # Return values in a fixed order so callers can unpack them even when # dataset wrappers add provenance fields like "dataset_id". - return [batch[key] for key in batch_keys] + [None, None] + return [batch[key] for key in batch_keys] + [None, packed_seq_params] # define spiky loss as a loss that's 10x the max loss observed @@ -287,6 +358,8 @@ def forward_step(data_iterator, model: HybridModel): with stimer(bdata=True): vp_stage = get_attr_wrapped_model(model, "vp_stage") + decoder = get_attr_wrapped_model(model, "decoder") + cp_partition_mode = decoder.cp_stage_entry_partition_mode ( attention_mask, cu_seqlens, @@ -300,34 +373,14 @@ def forward_step(data_iterator, model: HybridModel): tokens, padding_mask, packed_seq_params, - ) = get_batch(data_iterator, vp_stage) - - if packed_seq_params is not None: - if packed_seq_params.cu_seqlens_q is not None: - update_seqlen_stats_from_cu_seqlens(packed_seq_params.cu_seqlens_q) - elif cu_seqlens is not None: - # Squeeze the batch dim: the batch dict keeps cu_seqlens as (1, N) - # for consistency, but PackedSeqParams and TE expect 1-D. - cu_seqlens = cu_seqlens.squeeze(0) - if cu_seqlens_padded is not None: - cu_seqlens_padded = cu_seqlens_padded.squeeze(0) - # Use real (unpadded) cu_seqlens to feed the FLOPs accounting: varlen - # attention only computes work for real tokens within each chunk. + ) = get_batch(data_iterator, vp_stage, cp_partition_mode=cp_partition_mode) + + if cu_seqlens is not None: + if cu_seqlens.dim() == 2: + cu_seqlens = cu_seqlens.squeeze(0) update_seqlen_stats_from_cu_seqlens(cu_seqlens) - cu_seqlens_for_params = cu_seqlens_padded if cu_seqlens_padded is not None else cu_seqlens - packed_seq_params = PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu_seqlens_for_params, - cu_seqlens_kv=cu_seqlens_for_params, - cu_seqlens_q_padded=cu_seqlens_padded, - cu_seqlens_kv_padded=cu_seqlens_padded, - max_seqlen_q=int(max_seqlen.item()), - max_seqlen_kv=int(max_seqlen.item()), - local_cp_size=int(local_cp_size.item()) if local_cp_size is not None else None, - cp_group=hybrid_cp_group, - total_tokens=int(cu_seqlens_for_params[-1].item()), - tokens_per_sample=args.seq_length, - ) + elif packed_seq_params is not None and packed_seq_params.cu_seqlens_q is not None: + update_seqlen_stats_from_cu_seqlens(packed_seq_params.cu_seqlens_q) timers('batch-generator').stop() diff --git a/tests/unit_tests/models/test_gpt_model.py b/tests/unit_tests/models/test_gpt_model.py index 719b3781394..40c0afb62a5 100644 --- a/tests/unit_tests/models/test_gpt_model.py +++ b/tests/unit_tests/models/test_gpt_model.py @@ -3,6 +3,7 @@ import inspect import os from datetime import timedelta +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest @@ -12,6 +13,7 @@ from pytest import approx from transformer_engine.pytorch.fp8 import check_fp8_support +import megatron.core.models.gpt.gpt_model as gpt_model_module from megatron.core import parallel_state from megatron.core.hyper_comm_grid import HyperCommGrid from megatron.core.inference.config import InferenceConfig @@ -24,6 +26,7 @@ get_mlp_module_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.module import Float16Module @@ -32,6 +35,22 @@ from tests.unit_tests.test_utilities import Utils +class _FakeCpGroup: + def size(self): + return 2 + + +class _RecordingOutputLayer: + sequence_parallel = False + + def __init__(self): + self.hidden_states = None + + def __call__(self, hidden_states, **_kwargs): + self.hidden_states = hidden_states + return hidden_states, None + + class TestGPTModel: def setup_method(self, method): @@ -230,6 +249,69 @@ def output_processor(**kwargs): assert seen["output_layer"] is self.gpt_model.output_layer +def test_gpt_postprocess_restores_output_layout_to_input_layout(monkeypatch): + cp_group = _FakeCpGroup() + block_output_hidden = torch.arange(24, dtype=torch.float32).view(4, 2, 3) + input_layout_hidden = block_output_hidden + 1000.0 + output_layer = _RecordingOutputLayer() + calls = [] + + class FakeCpPartitionModeConverter: + def __init__(self, **kwargs): + self.kwargs = kwargs + + def convert(self, value, *, seq_dim=0, sequence_parallel=False): + calls.append((value, self.kwargs, seq_dim, sequence_parallel)) + assert value is block_output_hidden + assert self.kwargs["cp_group"] is cp_group + return input_layout_hidden + + monkeypatch.setattr( + gpt_model_module, "CpPartitionModeConverter", FakeCpPartitionModeConverter + ) + + model = object.__new__(GPTModel) + model.config = SimpleNamespace( + config_logger_dir="", + mtp_num_layers=0, + sequence_parallel=False, + use_mup=False, + ) + model.decoder = SimpleNamespace( + cp_stage_entry_partition_mode="zigzag", + ) + model.pg_collection = SimpleNamespace(cp=cp_group, tp=None) + model.share_embeddings_and_output_weights = False + model.post_process = True + model.output_layer = output_layer + model._scale_logits = lambda logits: logits + + output = GPTModel._postprocess( + model, + hidden_states=block_output_hidden, + input_ids=None, + position_ids=None, + labels=None, + rotary_pos_emb=None, + rotary_pos_cos=None, + rotary_pos_sin=None, + mtp_in_postprocess=False, + packed_seq_params=PackedSeqParams( + qkv_format="sbhd", + cp_group=cp_group, + cp_partition_mode="contiguous", + ), + ) + + assert output_layer.hidden_states is input_layout_hidden + assert torch.equal(output, input_layout_hidden.transpose(0, 1).contiguous()) + assert len(calls) == 1 + _, kwargs, seq_dim, _ = calls[0] + assert kwargs["source_partition_mode"] == "contiguous" + assert kwargs["target_partition_mode"] == "zigzag" + assert seq_dim == 0 + + def test_get_mlp_module_spec_interface(): # Get the function signature sig = inspect.signature(get_mlp_module_spec) diff --git a/tests/unit_tests/models/test_hybrid_model.py b/tests/unit_tests/models/test_hybrid_model.py index 6ac6525f30f..87a99745567 100644 --- a/tests/unit_tests/models/test_hybrid_model.py +++ b/tests/unit_tests/models/test_hybrid_model.py @@ -9,6 +9,7 @@ import torch from transformer_engine.pytorch.fp8 import check_fp8_support +import megatron.core.models.hybrid.hybrid_model as hybrid_model_module from megatron.core import parallel_state from megatron.core.hyper_comm_grid import HyperCommGrid from megatron.core.inference.config import InferenceConfig, MambaInferenceStateConfig @@ -73,6 +74,35 @@ def _get_dummy_hybrid_stack_spec() -> ModuleSpec: ) +class _FakeCpGroup: + def size(self): + return 2 + + +class _FakeHybridDecoder: + cp_stage_entry_partition_mode = "zigzag" + + def __init__(self, hidden_states): + self.hidden_states = hidden_states + + def __call__(self, **kwargs): + packed_seq_params = kwargs.get("packed_seq_params") + if packed_seq_params is not None: + packed_seq_params.cp_partition_mode = "contiguous" + return self.hidden_states + + +class _RecordingOutputLayer: + sequence_parallel = False + + def __init__(self): + self.hidden_states = None + + def __call__(self, hidden_states, **_kwargs): + self.hidden_states = hidden_states + return hidden_states, None + + def test_hybrid_logging_process_groups_are_paired(): tp_group = object() dp_cp_group = object() @@ -185,6 +215,71 @@ def test_hybrid_model_with_custom_process_groups(tmp_path, tp_size, cp_size, pp_ Utils.destroy_model_parallel() +def test_hybrid_forward_restores_output_layout_to_input_layout(monkeypatch): + cp_group = _FakeCpGroup() + block_output_hidden = torch.arange(24, dtype=torch.float32).view(4, 2, 3) + input_layout_hidden = block_output_hidden + 1000.0 + output_layer = _RecordingOutputLayer() + calls = [] + + class FakeCpPartitionModeConverter: + def __init__(self, **kwargs): + self.kwargs = kwargs + + def convert(self, value, *, seq_dim=0, sequence_parallel=False): + calls.append((value, self.kwargs, seq_dim, sequence_parallel)) + assert value is block_output_hidden + assert self.kwargs["cp_group"] is cp_group + return input_layout_hidden + + monkeypatch.setattr( + hybrid_model_module, "CpPartitionModeConverter", FakeCpPartitionModeConverter + ) + + model = object.__new__(HybridModel) + model.config = SimpleNamespace( + context_parallel_size=2, + cuda_graph_impl="none", + fine_grained_activation_offloading=False, + moe_n_hash_layers=0, + moe_paged_stash=False, + mtp_num_layers=0, + multi_latent_attention=False, + sequence_parallel=False, + use_mup=False, + ) + model.decoder = _FakeHybridDecoder(block_output_hidden) + model.pg_collection = SimpleNamespace(cp=cp_group, tp=None) + model.position_embedding_type = "none" + model.pre_process = False + model.post_process = True + model.mtp_process = False + model.share_embeddings_and_output_weights = False + model.output_layer = output_layer + model._scale_logits = lambda logits: logits + + output = HybridModel.forward( + model, + input_ids=torch.zeros(2, 4, dtype=torch.long), + position_ids=torch.zeros(2, 4, dtype=torch.long), + attention_mask=None, + decoder_input=torch.empty_like(block_output_hidden), + packed_seq_params=PackedSeqParams( + qkv_format="sbhd", + cp_group=cp_group, + cp_partition_mode="zigzag", + ), + ) + + assert output_layer.hidden_states is input_layout_hidden + assert torch.equal(output, input_layout_hidden.transpose(0, 1).contiguous()) + assert len(calls) == 1 + _, kwargs, seq_dim, _ = calls[0] + assert kwargs["source_partition_mode"] == "contiguous" + assert kwargs["target_partition_mode"] == "zigzag" + assert seq_dim == 0 + + class TestHybridModel: def setup_method(self, method): diff --git a/tests/unit_tests/ssm/test_gated_delta_net.py b/tests/unit_tests/ssm/test_gated_delta_net.py index 06e4c136d66..3d22f4e4a6d 100644 --- a/tests/unit_tests/ssm/test_gated_delta_net.py +++ b/tests/unit_tests/ssm/test_gated_delta_net.py @@ -14,6 +14,7 @@ ) from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( get_experimental_attention_variant_module_spec, + get_experimental_attention_variant_stage_input_cp_partition_mode, get_transformer_block_with_experimental_attention_variant_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel @@ -100,6 +101,24 @@ def _make_gdn_config(**overrides): return TransformerConfig(**config_kwargs) +def _set_gdn_test_cp_partition_mode(packed_seq_params, cp_size, linear_cp_mode): + if cp_size <= 1: + return packed_seq_params + if linear_cp_mode == "headwise": + packed_seq_params.cp_partition_mode = "zigzag" + elif linear_cp_mode == "chunkwise": + packed_seq_params.cp_partition_mode = "contiguous" + else: + raise ValueError(f"Invalid linear CP mode: {linear_cp_mode}") + return packed_seq_params + + +def _make_sbhd_cp_packed_seq_params(cp_group, cp_partition_mode): + return PackedSeqParams( + qkv_format="sbhd", cp_group=cp_group, cp_partition_mode=cp_partition_mode + ) + + def test_gdn_pre_gated_delta_rule_fusion_defaults_to_disabled(): config = _make_gdn_config() assert not config.gdn_pre_gated_delta_rule_fusion @@ -543,6 +562,7 @@ def test_gpu_forward_thd_correctness(self): hidden_states_thd = hidden_states_thd.view(-1, 1, self.gdn.config.hidden_size) attention_mask_thd = None packed_seq_params = make_test_packed_seq_params(cu_seqlens=cu_seqlens) + _set_gdn_test_cp_partition_mode(packed_seq_params, self.cp_size, self.linear_cp_mode) # THD format output_thd, _ = self.gdn( @@ -594,6 +614,7 @@ def test_gpu_forward_thd_padding_correctness(self): padded_params = make_test_packed_seq_params_with_padding( cu_seqlens=[0, 30, 60, 90, 120], cu_seqlens_padded=[0, 32, 64, 96, 128] ) + _set_gdn_test_cp_partition_mode(padded_params, self.cp_size, self.linear_cp_mode) output_thd_padded, _ = self.gdn(hidden_states_thd, None, packed_seq_params=padded_params) output_thd2bshd = output_thd_padded.view(*output_bshd.shape) torch.testing.assert_close( @@ -606,6 +627,7 @@ def test_gpu_forward_thd_padding_correctness(self): # B) no-padded branch: use actual cu_seqlens when it matches total_sequence_length. no_padding_params = make_test_packed_seq_params(cu_seqlens=[0, 32, 64, 96, 128]) + _set_gdn_test_cp_partition_mode(no_padding_params, self.cp_size, self.linear_cp_mode) output_thd_no_padding, _ = self.gdn( hidden_states_thd, None, packed_seq_params=no_padding_params ) @@ -631,11 +653,13 @@ def test_gpu_forward_thd_padding_correctness(self): padded_mismatch_params = make_test_packed_seq_params_with_padding( cu_seqlens=[0, 30, 60, 90, 120], cu_seqlens_padded=[0, 32, 64, 96, 126] ) + _set_gdn_test_cp_partition_mode(padded_mismatch_params, self.cp_size, self.linear_cp_mode) with pytest.raises(ValueError, match="does not match"): self.gdn(hidden_states_thd, None, packed_seq_params=padded_mismatch_params) # E) actual mismatch branch without *_padded: should raise. actual_mismatch_params = make_test_packed_seq_params(cu_seqlens=[0, 32, 64, 96, 129]) + _set_gdn_test_cp_partition_mode(actual_mismatch_params, self.cp_size, self.linear_cp_mode) with pytest.raises(ValueError, match="does not match"): self.gdn(hidden_states_thd, None, packed_seq_params=actual_mismatch_params) @@ -748,6 +772,15 @@ def _assert_pre_gated_delta_rule_outputs_close( msg=lambda msg, output_name=name: f"{output_name} mismatch: {msg}", ) + def _make_pre_gated_delta_rule_grad_outputs(self, outputs): + grad_outputs = [] + for output_idx, output in enumerate(outputs): + grad = torch.linspace( + -0.1, 0.1, output.numel(), device=output.device, dtype=torch.float32 + ).reshape(output.shape) + grad_outputs.append(grad + (output_idx - 2.5) * 0.01) + return grad_outputs + def test_fused_and_unfused_forward_match(self): hidden_states = torch.randn( (32, 2, self.unfused_gdn.config.hidden_size), @@ -941,6 +974,7 @@ def test_fused_and_unfused_pre_gated_delta_rule_backward_match(self): batch = 2 seq_len = 32 + torch.manual_seed(1234) qkvzba = torch.randn( (seq_len, batch, reference_gdn.in_proj_dim), device=torch.cuda.current_device(), @@ -956,7 +990,7 @@ def test_fused_and_unfused_pre_gated_delta_rule_backward_match(self): qkvzba_unfused, batch, seq_len, reference_gdn.cp_size, reference_gdn.pg_collection.cp ) fused_outputs = fused_gdn._fused_streamed_pre_gated_delta_rule(qkvzba_fused) - grad_outputs = [torch.randn_like(output.float()) for output in unfused_outputs] + grad_outputs = self._make_pre_gated_delta_rule_grad_outputs(unfused_outputs) unfused_loss = sum( (output.float() * grad).sum() for output, grad in zip(unfused_outputs, grad_outputs) @@ -992,6 +1026,7 @@ def test_fused_and_unfused_packed_pre_gated_delta_rule_forward_match(self): [0, 1, 4, 6, 11], device=torch.cuda.current_device(), dtype=torch.int32 ) seq_len = cu_seqlens[-1].item() + torch.manual_seed(1234) qkvzba = torch.randn( (seq_len, batch, reference_gdn.in_proj_dim), device=torch.cuda.current_device(), @@ -1095,7 +1130,7 @@ def test_fused_and_unfused_packed_pre_gated_delta_rule_backward_match(self): fused_outputs = fused_gdn._fused_streamed_pre_gated_delta_rule( qkvzba_fused, cu_seqlens_q=cu_seqlens ) - grad_outputs = [torch.randn_like(output.float()) for output in unfused_outputs] + grad_outputs = self._make_pre_gated_delta_rule_grad_outputs(unfused_outputs) unfused_loss = sum( (output.float() * grad).sum() for output, grad in zip(unfused_outputs, grad_outputs) @@ -1401,6 +1436,7 @@ def _make_packed_seq_params(cu_seqlens): cu_seqlens[i + 1] - cu_seqlens[i] for i in range(len(cu_seqlens) - 1) ), total_tokens=cu_seqlens[-1] // cp_size, + cp_partition_mode="contiguous", ) @staticmethod @@ -1734,7 +1770,205 @@ def test_parallel_gated_delta_net_correctness( sequence_length=256, micro_batch_size=micro_batch_size, sequence_packing=sequence_packing, + cp_partition_mode="contiguous" if is_chunkwise_cp else "zigzag", + cp_stage_entry_partition_mode="contiguous" if is_chunkwise_cp else "zigzag", + compare_param_grads=is_chunkwise_cp and tp == 1 and not sequence_packing, + ) + + +@pytest.mark.skipif(not HAVE_FLA, reason="FLA is not installed.") +@pytest.mark.internal +@pytest.mark.parametrize("cp", [2, 4]) +def test_mixed_gdn_sdpa_gpt_model_cp_boundary_forward_backward_correctness(tmp_path_dist_ckpt, cp): + if not torch.cuda.is_available() or Utils.world_size < cp: + pytest.skip(f"Mixed GDN/SDPA CP parity needs at least {cp} CUDA ranks.") + + sequence_length = 64 + micro_batch_size = 1 + vocab_size = 128 + seed = 123 + + def make_config(context_parallel_size): + return TransformerConfig( + hidden_size=128, + linear_conv_kernel_dim=2, + linear_key_head_dim=32, + linear_value_head_dim=32, + linear_num_key_heads=4, + linear_num_value_heads=8, + num_layers=4, + normalization="RMSNorm", + use_cpu_initialization=True, + layernorm_zero_centered_gamma=True, + num_attention_heads=8, + activation_func=F.silu, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=4, + linear_cp_mode="chunkwise", + transformer_impl="transformer_engine", + context_parallel_size=context_parallel_size, + hidden_dropout=0.0, + attention_dropout=0.0, + bf16=True, + params_dtype=torch.bfloat16, + ) + + def initialize_gpt_model( + config, pre_process=True, post_process=True, vp_stage=None, pg_collection=None + ): + transformer_layer_spec = get_transformer_block_with_experimental_attention_variant_spec( + config=config, vp_stage=vp_stage, pp_rank=0 + ) + cp_stage_entry_partition_mode = ( + get_experimental_attention_variant_stage_input_cp_partition_mode(config) + ) + return GPTModel( + config=config, + transformer_layer_spec=transformer_layer_spec, + vocab_size=vocab_size, + max_sequence_length=sequence_length, + pre_process=pre_process, + post_process=post_process, + position_embedding_type="rope", + pg_collection=pg_collection, + vp_stage=vp_stage, + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, + ) + + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, context_parallel_size=1 + ) + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + input_ids = torch.randint( + low=0, + high=vocab_size, + size=(micro_batch_size, sequence_length), + device=torch.device(f"cuda:{torch.cuda.current_device()}"), ) + position_ids = torch.arange( + sequence_length, device=input_ids.device, dtype=torch.long + ).unsqueeze(0) + labels = (input_ids + 1) % vocab_size + + def _get_param_grad(param): + grad = param.grad + if grad is None: + grad = getattr(param, "main_grad", None) + if grad is None: + param.grad = torch.zeros_like(param) + grad = param.grad + return grad + + def _scale_grads(model, scale): + for param in model.parameters(): + if param.requires_grad: + _get_param_grad(param).data.mul_(scale) + + def _all_reduce_grads(model, group): + for param in model.parameters(): + if param.requires_grad: + torch.distributed.all_reduce(_get_param_grad(param), group=group) + + def _collect_grads(model): + return { + name: _get_param_grad(param).detach().float().clone() + for name, param in model.named_parameters() + if param.requires_grad + } + + def _zero_grads(model): + for param in model.parameters(): + param.grad = None + main_grad = getattr(param, "main_grad", None) + if main_grad is not None: + main_grad.zero_() + + with TempNamedDir(tmp_path_dist_ckpt / 'test_mixed_gdn_sdpa_gpt_cp', sync=True) as ckpt_dir: + mock_args = parse_args(ignore_unknown_args=True) + set_args(mock_args) + + baseline_config = make_config(context_parallel_size=1) + init_basic_mock_args(mock_args, 1, 1, bf16=True) + mock_args.context_parallel_size = 1 + baseline_model = unwrap_model(get_model(initialize_gpt_model, config=baseline_config)) + baseline_model[0].eval() + + init_checkpointing_mock_args(mock_args, ckpt_dir, False) + mock_args.no_save_optim = True + mock_args.no_save_rng = True + mock_args.no_load_optim = True + mock_args.no_load_rng = True + save_checkpoint(10, baseline_model, None, None, 0) + + _zero_grads(baseline_model[0]) + baseline_loss = baseline_model[0]( + input_ids=input_ids, position_ids=position_ids, attention_mask=None, labels=labels + ) + baseline_loss.float().sum().backward() + _scale_grads(baseline_model[0], 1.0 / baseline_loss.numel()) + baseline_grads = _collect_grads(baseline_model[0]) + baseline_loss = baseline_loss.detach() + + Utils.destroy_model_parallel() + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, context_parallel_size=cp + ) + model_parallel_cuda_manual_seed(seed) + parallel_config = make_config(context_parallel_size=cp) + init_basic_mock_args(mock_args, 1, 1, bf16=True) + mock_args.context_parallel_size = cp + parallel_model = unwrap_model(get_model(initialize_gpt_model, config=parallel_config)) + parallel_model[0].eval() + with mock.patch('megatron.training.checkpointing.check_checkpoint_args'): + with mock.patch('megatron.training.checkpointing.update_num_microbatches'): + load_checkpoint(parallel_model, None, None) + + cp_group = parallel_state.get_context_parallel_group() + input_partition_mode = parallel_model[0].decoder.cp_stage_entry_partition_mode + assert input_partition_mode == "contiguous" + + local_input_ids = get_tensor_on_this_cp_rank( + input_ids, 1, cp_group, cp_partition_mode=input_partition_mode + ) + local_position_ids = get_tensor_on_this_cp_rank( + position_ids, 1, cp_group, cp_partition_mode=input_partition_mode + ) + local_labels = get_tensor_on_this_cp_rank( + labels, 1, cp_group, cp_partition_mode=input_partition_mode + ) + _zero_grads(parallel_model[0]) + parallel_loss = parallel_model[0]( + input_ids=local_input_ids, + position_ids=local_position_ids, + attention_mask=None, + labels=local_labels, + packed_seq_params=_make_sbhd_cp_packed_seq_params(cp_group, input_partition_mode), + ) + + expected_loss = get_tensor_on_this_cp_rank( + baseline_loss, 1, cp_group, cp_partition_mode=input_partition_mode + ) + torch.testing.assert_close( + parallel_loss.float(), expected_loss.float(), atol=2e-3, rtol=2e-3 + ) + + parallel_loss.float().sum().backward() + _all_reduce_grads(parallel_model[0], cp_group) + _scale_grads(parallel_model[0], 1.0 / baseline_loss.numel()) + parallel_grads = _collect_grads(parallel_model[0]) + + assert baseline_grads.keys() == parallel_grads.keys() + for name, baseline_grad in baseline_grads.items(): + torch.testing.assert_close( + parallel_grads[name], + baseline_grad, + atol=2e-3, + rtol=2e-3, + msg=lambda msg, param_name=name: f"gradient mismatch for {param_name}: {msg}", + ) + + Utils.destroy_model_parallel() @pytest.mark.parametrize("cp_size", [2, 4], scope="class") diff --git a/tests/unit_tests/ssm/test_gdn_moe_cp_loss_parity.py b/tests/unit_tests/ssm/test_gdn_moe_cp_loss_parity.py new file mode 100644 index 00000000000..8cb1dbb50f7 --- /dev/null +++ b/tests/unit_tests/ssm/test_gdn_moe_cp_loss_parity.py @@ -0,0 +1,641 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import warnings +from copy import deepcopy +from inspect import signature + +import pytest +import torch +import torch.nn.functional as F + +from megatron.core import parallel_state +from megatron.core.context_parallel_layout import prebuild_thd_cp_partition_routes +from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_stage_input_cp_partition_mode, + get_transformer_block_with_experimental_attention_variant_spec, +) +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec +from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.optimizer.clip_grads import get_grad_norm_fp32 +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer import TransformerConfig +from megatron.core.transformer.multi_token_prediction import ( + MTPLossAutoScaler, + MTPLossLoggingHelper, +) +from megatron.core.utils import ( + flatten_batch_for_packed_sequences, + get_batch_on_this_cp_rank, + get_thd_batch_on_this_cp_rank, + is_te_min_version, +) +from megatron.training.arguments import parse_args +from megatron.training.global_vars import set_args +from tests.unit_tests.dist_checkpointing import init_basic_mock_args +from tests.unit_tests.test_utilities import Utils + +try: + import fla # noqa: F401 + + HAVE_FLA = True +except ImportError: + HAVE_FLA = False + + +_PARALLEL_EP_SIZE = 4 +_SEQUENCE_LENGTH = 4096 +_MICRO_BATCH_SIZE = 1 +_VOCAB_SIZE = 8192 +_SEED = 1234 +_DIAGNOSTIC_REPEATS = 5 +_NUM_LAYERS = 3 +_LINEAR_ATTENTION_PATTERN = [1, 0, 0] +_SIDE_LAYOUT_GUARD_LAYER_INDEX = 2 +_SIDE_LAYOUT_GUARD_PARTITION_MODE = "zigzag" +_LM_LOSS_ATOL = 0.003 +_MTP_LOSS_ATOL = 0.003 +_GRAD_NORM_ATOL = 0.1 +_GRAD_NORM_RTOL = 0.01 +_REFERENCE_LINEAR_CP_MODE = "headwise" +_CANDIDATE_LINEAR_CP_MODE = "chunkwise" +_CP_LAYOUT_WARNING_SUBSTRINGS = ( + "missing precomputed context-parallel layout routes", + "received no PackedSeqParams while running under context parallelism", +) +_PARALLEL_CASES = ( + pytest.param(4, 1, _PARALLEL_EP_SIZE, False, id="cp4_tp1_ep4"), + pytest.param(2, 2, _PARALLEL_EP_SIZE, True, id="cp2_tp2_ep4_sp"), +) +_FULL_RECOMPUTE_CASES = ( + pytest.param(False, id="no_recompute"), + pytest.param(True, id="full_recompute"), +) + + +def _destroy_model_parallel_without_barrier(): + if not Utils.inited: + return + torch.cuda.synchronize() + parallel_state.destroy_model_parallel() + Utils.inited = False + torch.cuda.memory.empty_cache() + + +def _collect_parameter_state(model): + return {name: param.detach().cpu().clone() for name, param in model.named_parameters()} + + +def _copy_state_to_model(source_state, model): + with torch.no_grad(): + for name, param in model.named_parameters(): + source = source_state[name] + if source.shape == param.shape: + param.copy_(source.to(device=param.device, dtype=param.dtype)) + continue + + raise AssertionError( + f"Cannot copy parameter {name}: source shape {tuple(source.shape)}, " + f"target shape {tuple(param.shape)}" + ) + + +def _install_layer_rotary_layout_guard( + model, layer_index, expected_rotary_pos_emb, expected_partition_mode +): + """Assert a decoder layer receives RoPE in the expected CP layout.""" + layer = model.decoder.layers[layer_index] + original_forward = layer.forward + + def guarded_forward(*args, **kwargs): + packed_seq_params = kwargs.get("packed_seq_params") + actual_partition_mode = getattr(packed_seq_params, "cp_partition_mode", None) + assert actual_partition_mode == expected_partition_mode, ( + f"Layer {layer_index} expected packed_seq_params.cp_partition_mode=" + f"{expected_partition_mode!r}, got {actual_partition_mode!r}." + ) + rotary_pos_emb = kwargs.get("rotary_pos_emb") + assert rotary_pos_emb is not None, f"Layer {layer_index} did not receive RoPE." + torch.testing.assert_close( + rotary_pos_emb, + expected_rotary_pos_emb, + atol=0.0, + rtol=0.0, + msg=( + f"Layer {layer_index} received RoPE that does not match " + f"{expected_partition_mode!r} CP layout." + ), + ) + return original_forward(*args, **kwargs) + + layer.forward = guarded_forward + + +def _assert_no_cp_layout_warnings(caught_warnings): + for caught_warning in caught_warnings: + message = str(caught_warning.message) + assert not any( + warning_substring in message + for warning_substring in _CP_LAYOUT_WARNING_SUBSTRINGS + ), message + + +def _get_model_input_partition_mode(model): + decoder = getattr(model, "decoder", None) + if decoder is None: + return "zigzag" + return decoder.cp_stage_entry_partition_mode or "zigzag" + + +def _get_stage_entry_partition_mode(config, vp_stage=None, pp_rank=0): + return get_experimental_attention_variant_stage_input_cp_partition_mode( + config=config, vp_stage=vp_stage, pp_rank=pp_rank + ) + + +def _make_config( + context_parallel_size, + tensor_model_parallel_size, + expert_model_parallel_size, + sequence_parallel, + qkv_format, + linear_cp_mode, + full_recompute, +): + packed_kwargs = {} + if qkv_format == "thd": + packed_kwargs = { + "sequence_packing_scheduler": "dp_balanced", + "pad_packed_seq_alignment": "max", + "max_seqlen_per_dp_cp_rank": _SEQUENCE_LENGTH, + } + + recompute_kwargs = ( + { + "recompute_granularity": "full", + "recompute_method": "uniform", + "recompute_num_layers": 1, + } + if full_recompute + else { + "recompute_granularity": None, + "recompute_method": None, + "recompute_num_layers": None, + } + ) + + return TransformerConfig( + hidden_size=512, + ffn_hidden_size=1024, + linear_conv_kernel_dim=4, + linear_key_head_dim=64, + linear_value_head_dim=64, + linear_num_key_heads=4, + linear_num_value_heads=8, + num_layers=_NUM_LAYERS, + normalization="RMSNorm", + layernorm_epsilon=1e-6, + use_cpu_initialization=True, + layernorm_zero_centered_gamma=True, + num_attention_heads=8, + kv_channels=64, + num_query_groups=2, + qk_layernorm=True, + attention_output_gate=True, + activation_func=F.silu, + gated_linear_unit=True, + add_bias_linear=False, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=_LINEAR_ATTENTION_PATTERN, + linear_cp_mode=linear_cp_mode, + transformer_impl="transformer_engine", + tensor_model_parallel_size=tensor_model_parallel_size, + expert_model_parallel_size=expert_model_parallel_size, + expert_tensor_parallel_size=1, + context_parallel_size=context_parallel_size, + cp_partition_mode="auto", + sequence_parallel=sequence_parallel, + hidden_dropout=0.0, + attention_dropout=0.0, + calculate_per_token_loss=True, + bf16=True, + params_dtype=torch.bfloat16, + num_moe_experts=32, + moe_layer_freq=1, + moe_ffn_hidden_size=128, + moe_shared_expert_intermediate_size=128, + moe_shared_expert_gate=True, + moe_router_load_balancing_type="aux_loss", + moe_router_topk=4, + moe_grouped_gemm=True, + moe_aux_loss_coeff=0.0, + moe_token_dispatcher_type="flex", + moe_flex_dispatcher_backend="hybridep", + moe_flex_dispatcher_num_sms=32, + moe_permute_fusion=True, + moe_router_fusion=True, + moe_router_dtype="fp32", + mtp_num_layers=1, + mtp_loss_scaling_factor=1.0, + mtp_use_repeated_layer=False, + gdn_pre_gated_delta_rule_fusion=False, + **packed_kwargs, + **recompute_kwargs, + ) + + +def _initialize_gpt_model( + config, pre_process=True, post_process=True, vp_stage=None, pg_collection=None +): + transformer_layer_spec = get_transformer_block_with_experimental_attention_variant_spec( + config=config, vp_stage=vp_stage, pp_rank=0 + ) + mtp_block_spec = None + if config.mtp_num_layers: + mtp_block_spec = get_gpt_mtp_block_spec( + config=config, + spec=transformer_layer_spec, + use_transformer_engine=True, + vp_stage=vp_stage, + pp_rank=0, + ) + + model_kwargs = { + "config": config, + "transformer_layer_spec": transformer_layer_spec, + "mtp_block_spec": mtp_block_spec, + "vocab_size": _VOCAB_SIZE, + "max_sequence_length": _SEQUENCE_LENGTH, + "pre_process": pre_process, + "post_process": post_process, + "position_embedding_type": "rope", + "rotary_percent": 0.25, + "rotary_base": 10000000, + "pg_collection": pg_collection, + "vp_stage": vp_stage, + } + if "cp_stage_entry_partition_mode" in signature(GPTModel).parameters: + model_kwargs["cp_stage_entry_partition_mode"] = _get_stage_entry_partition_mode( + config=config, vp_stage=vp_stage, pp_rank=0 + ) + + return GPTModel(**model_kwargs) + + +def _build_gpt_model(config, device): + model = _initialize_gpt_model(config) + model.to(device=device) + return model + + +def _set_mock_args(args, config, context_parallel_size): + init_basic_mock_args(args, config.tensor_model_parallel_size, 1, bf16=True) + args.context_parallel_size = context_parallel_size + args.cp_comm_type = "a2a" if context_parallel_size == 1 else "p2p" + args.expert_model_parallel_size = config.expert_model_parallel_size + args.expert_tensor_parallel_size = 1 + args.num_experts = config.num_moe_experts + args.moe_ffn_hidden_size = config.moe_ffn_hidden_size + args.moe_shared_expert_intermediate_size = config.moe_shared_expert_intermediate_size + args.moe_shared_expert_gate = config.moe_shared_expert_gate + args.moe_router_load_balancing_type = config.moe_router_load_balancing_type + args.moe_router_topk = config.moe_router_topk + args.moe_grouped_gemm = config.moe_grouped_gemm + args.moe_aux_loss_coeff = config.moe_aux_loss_coeff + args.moe_token_dispatcher_type = config.moe_token_dispatcher_type + args.moe_flex_dispatcher_backend = config.moe_flex_dispatcher_backend + args.moe_flex_dispatcher_num_sms = config.moe_flex_dispatcher_num_sms + args.moe_permute_fusion = config.moe_permute_fusion + args.moe_router_fusion = config.moe_router_fusion + args.moe_router_dtype = config.moe_router_dtype + args.mtp_num_layers = config.mtp_num_layers + args.mtp_loss_scaling_factor = config.mtp_loss_scaling_factor + args.mtp_use_repeated_layer = config.mtp_use_repeated_layer + args.recompute_granularity = config.recompute_granularity + args.recompute_method = config.recompute_method + args.recompute_num_layers = config.recompute_num_layers + args.linear_cp_mode = config.linear_cp_mode + args.sequence_parallel = config.sequence_parallel + args.seq_length = _SEQUENCE_LENGTH + args.max_position_embeddings = _SEQUENCE_LENGTH + args.padded_vocab_size = _VOCAB_SIZE + args.untie_embeddings_and_output_weights = True + + +def _make_sbhd_batch(device): + tokens = torch.randint( + low=0, + high=_VOCAB_SIZE, + size=(_MICRO_BATCH_SIZE, _SEQUENCE_LENGTH), + device=device, + dtype=torch.long, + ) + valid_length = 3512 + prompt_length = 67 + tokens[:, valid_length:] = 0 + labels = (tokens + 1) % _VOCAB_SIZE + loss_mask = torch.zeros_like(tokens, dtype=torch.float32) + loss_mask[:, prompt_length:valid_length] = 1.0 + position_ids = torch.arange(_SEQUENCE_LENGTH, device=device, dtype=torch.long).unsqueeze(0) + return { + "tokens": tokens, + "labels": labels, + "loss_mask": loss_mask, + "attention_mask": None, + "position_ids": position_ids, + } + + +def _make_thd_batch(device): + padded_seq_lengths = [1024, 768, 1280, 1024] + seq_lengths = [901, 629, 1103, 877] + prompt_lengths = [0, 63, 257, 15] + assert sum(padded_seq_lengths) == _SEQUENCE_LENGTH + + tokens = torch.zeros((_MICRO_BATCH_SIZE, _SEQUENCE_LENGTH), device=device, dtype=torch.long) + labels = torch.zeros_like(tokens) + loss_mask = torch.zeros_like(tokens, dtype=torch.float32) + padding_mask = torch.ones_like(tokens, dtype=torch.bool) + position_ids = torch.empty_like(tokens) + + padded_offset = 0 + cu_seqlens = [0] + cu_seqlens_padded = [0] + for seq_length, padded_seq_length, prompt_length in zip( + seq_lengths, padded_seq_lengths, prompt_lengths + ): + valid_end = padded_offset + seq_length + padded_end = padded_offset + padded_seq_length + seq_tokens = torch.randint( + low=0, + high=_VOCAB_SIZE, + size=(_MICRO_BATCH_SIZE, seq_length), + device=device, + dtype=torch.long, + ) + tokens[:, padded_offset:valid_end] = seq_tokens + labels[:, padded_offset:valid_end] = (seq_tokens + 1) % _VOCAB_SIZE + loss_mask[:, padded_offset + prompt_length : valid_end] = 1.0 + padding_mask[:, padded_offset:valid_end] = False + position_ids[:, padded_offset:valid_end] = torch.arange( + seq_length, device=device, dtype=torch.long + ) + position_ids[:, valid_end:padded_end] = 0 + cu_seqlens.append(cu_seqlens[-1] + seq_length) + cu_seqlens_padded.append(padded_end) + padded_offset = padded_end + + cu_seqlens = torch.tensor([cu_seqlens], device=device, dtype=torch.int32) + cu_seqlens_padded = torch.tensor([cu_seqlens_padded], device=device, dtype=torch.int32) + max_seqlen = torch.tensor([max(padded_seq_lengths)], device=device, dtype=torch.int32) + return { + "tokens": tokens, + "labels": labels, + "loss_mask": loss_mask, + "attention_mask": None, + "padding_mask": padding_mask, + "position_ids": position_ids, + "cu_seqlens": cu_seqlens, + "cu_seqlens_padded": cu_seqlens_padded, + "max_seqlen": max_seqlen, + } + + +def _prepare_batch_for_model(batch, qkv_format, cp_group, cp_partition_mode): + batch = deepcopy(batch) + if qkv_format == "sbhd": + batch = get_batch_on_this_cp_rank( + batch, + is_hybrid_cp=False, + cp_group=cp_group, + cp_partition_mode=cp_partition_mode, + ) + packed_seq_params = PackedSeqParams( + qkv_format="sbhd", + cp_group=cp_group, + cp_partition_mode=cp_partition_mode, + ) + return batch, packed_seq_params + + batch = flatten_batch_for_packed_sequences(batch) + batch, packed_seq_params = get_thd_batch_on_this_cp_rank( + batch, + batch["cu_seqlens"][0], + batch["cu_seqlens_padded"][0], + batch["max_seqlen"], + cp_partition_mode=cp_partition_mode, + ) + prebuild_thd_cp_partition_routes(packed_seq_params, cp_group) + return batch, packed_seq_params + + +def _global_grad_norm(model): + grads = [ + param.grad.detach() + for param in model.parameters() + if param.requires_grad and param.grad is not None + ] + return torch.tensor( + get_grad_norm_fp32(grads, grad_stats_parallel_group=torch.distributed.group.WORLD), + device=torch.cuda.current_device(), + dtype=torch.float32, + ) + + +def _get_mtp_losses_from_tracker(): + MTPLossLoggingHelper.reduce_loss_in_tracker() + tracker = MTPLossLoggingHelper.tracker + assert "values" in tracker, "MTP loss tracker did not record any loss values." + return tracker["values"].detach().float().clone() + + +def _loss_and_grad_stats(model, batch, packed_seq_params): + model.zero_grad(set_to_none=True) + MTPLossLoggingHelper.clean_loss_in_tracker() + MTPLossAutoScaler.set_loss_scale(torch.ones((), device=torch.cuda.current_device())) + + with warnings.catch_warnings(record=True) as caught_warnings: + warnings.simplefilter("always") + loss = model( + input_ids=batch["tokens"], + position_ids=batch["position_ids"], + attention_mask=batch["attention_mask"], + labels=batch["labels"], + loss_mask=batch["loss_mask"], + packed_seq_params=packed_seq_params, + padding_mask=batch.get("padding_mask"), + ) + mtp_losses = _get_mtp_losses_from_tracker() + numerator = (loss.float() * batch["loss_mask"]).sum() + denominator = batch["loss_mask"].sum() + (numerator / denominator.clamp(min=1)).backward() + + _assert_no_cp_layout_warnings(caught_warnings) + grad_norm = _global_grad_norm(model) + return numerator, denominator, mtp_losses, grad_norm + + +@pytest.mark.internal +@pytest.mark.skipif(not HAVE_FLA, reason="FLA is not installed.") +@pytest.mark.skipif(not is_te_min_version("1.11.0"), reason="MoE grouped GEMM requires TE >= 1.11.") +@pytest.mark.parametrize("repeat_index", range(_DIAGNOSTIC_REPEATS)) +@pytest.mark.parametrize("full_recompute", _FULL_RECOMPUTE_CASES) +@pytest.mark.parametrize("qkv_format", ["sbhd", "thd"]) +@pytest.mark.parametrize( + ( + "context_parallel_size,tensor_model_parallel_size,expert_model_parallel_size," + "sequence_parallel" + ), + _PARALLEL_CASES, +) +def test_qwen35_proxy_gdn_moe_chunkwise_loss_and_grad_matches_headwise( + qkv_format, + context_parallel_size, + tensor_model_parallel_size, + expert_model_parallel_size, + sequence_parallel, + full_recompute, + repeat_index, +): + min_world_size = max( + context_parallel_size * tensor_model_parallel_size, expert_model_parallel_size + ) + if not torch.cuda.is_available() or Utils.world_size < min_world_size: + pytest.skip(f"GDN/MoE CP loss parity needs at least {min_world_size} CUDA ranks.") + + mock_args = parse_args(ignore_unknown_args=True) + set_args(mock_args) + + try: + Utils.initialize_model_parallel( + tensor_model_parallel_size=tensor_model_parallel_size, + pipeline_model_parallel_size=1, + expert_model_parallel_size=expert_model_parallel_size, + expert_tensor_parallel_size=1, + context_parallel_size=context_parallel_size, + ) + device = torch.device(f"cuda:{torch.cuda.current_device()}") + seed = _SEED + repeat_index + torch.manual_seed(seed) + batch = _make_thd_batch(device) if qkv_format == "thd" else _make_sbhd_batch(device) + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + reference_config = _make_config( + context_parallel_size=context_parallel_size, + tensor_model_parallel_size=tensor_model_parallel_size, + expert_model_parallel_size=expert_model_parallel_size, + sequence_parallel=sequence_parallel, + qkv_format=qkv_format, + linear_cp_mode=_REFERENCE_LINEAR_CP_MODE, + full_recompute=full_recompute, + ) + _set_mock_args(mock_args, reference_config, context_parallel_size=context_parallel_size) + + reference_model = _build_gpt_model(reference_config, device) + reference_model.train() + + cp_group = parallel_state.get_context_parallel_group() + reference_input_partition_mode = _get_model_input_partition_mode(reference_model) + reference_batch, reference_packed_seq_params = _prepare_batch_for_model( + batch, + qkv_format=qkv_format, + cp_group=cp_group, + cp_partition_mode=reference_input_partition_mode, + ) + reference_num, reference_den, reference_mtp_losses, reference_grad_norm = ( + _loss_and_grad_stats(reference_model, reference_batch, reference_packed_seq_params) + ) + reference_stats = torch.stack([reference_num.detach(), reference_den.detach()]) + torch.distributed.all_reduce(reference_stats, group=cp_group) + reference_avg = reference_stats[0] / reference_stats[1].clamp(min=1) + source_state = _collect_parameter_state(reference_model) + + del reference_model + torch.cuda.empty_cache() + + model_parallel_cuda_manual_seed(seed) + candidate_config = _make_config( + context_parallel_size=context_parallel_size, + tensor_model_parallel_size=tensor_model_parallel_size, + expert_model_parallel_size=expert_model_parallel_size, + sequence_parallel=sequence_parallel, + qkv_format=qkv_format, + linear_cp_mode=_CANDIDATE_LINEAR_CP_MODE, + full_recompute=full_recompute, + ) + _set_mock_args(mock_args, candidate_config, context_parallel_size=context_parallel_size) + candidate_model = _build_gpt_model(candidate_config, device) + candidate_model.train() + _copy_state_to_model(source_state, candidate_model) + + candidate_input_partition_mode = _get_model_input_partition_mode(candidate_model) + candidate_batch, candidate_packed_seq_params = _prepare_batch_for_model( + batch, + qkv_format=qkv_format, + cp_group=cp_group, + cp_partition_mode=candidate_input_partition_mode, + ) + if qkv_format == "sbhd": + expected_rotary_pos_emb = candidate_model.rotary_pos_emb( + _SEQUENCE_LENGTH, + packed_seq=False, + cp_group=cp_group, + cp_partition_mode=_SIDE_LAYOUT_GUARD_PARTITION_MODE, + ) + _install_layer_rotary_layout_guard( + candidate_model, + layer_index=_SIDE_LAYOUT_GUARD_LAYER_INDEX, + expected_rotary_pos_emb=expected_rotary_pos_emb, + expected_partition_mode=_SIDE_LAYOUT_GUARD_PARTITION_MODE, + ) + candidate_num, candidate_den, candidate_mtp_losses, candidate_grad_norm = ( + _loss_and_grad_stats(candidate_model, candidate_batch, candidate_packed_seq_params) + ) + stats = torch.stack([candidate_num.detach(), candidate_den.detach()]) + torch.distributed.all_reduce(stats, group=cp_group) + candidate_avg = stats[0] / stats[1].clamp(min=1) + lm_loss_diff = (candidate_avg.float() - reference_avg.float()).abs() + mtp_loss_diff = (candidate_mtp_losses - reference_mtp_losses).abs().max() + grad_norm_diff = (candidate_grad_norm - reference_grad_norm).abs() + + if torch.distributed.get_rank() == 0: + print( + "GDN MoE CP loss/grad parity: " + f"case=cp{context_parallel_size}_tp{tensor_model_parallel_size}_" + f"ep{expert_model_parallel_size}" + f"{'_sp' if sequence_parallel else ''} " + f"layers={_NUM_LAYERS} linear_attention_pattern={_LINEAR_ATTENTION_PATTERN} " + f"format={qkv_format} full_recompute={full_recompute} " + f"repeat={repeat_index} seed={seed} " + f"lm_reference={reference_avg.float().item():.8f} " + f"lm_candidate={candidate_avg.float().item():.8f} " + f"lm_diff={lm_loss_diff.item():.8f} " + f"mtp_reference={reference_mtp_losses[0].item():.8f} " + f"mtp_candidate={candidate_mtp_losses[0].item():.8f} " + f"mtp_diff={mtp_loss_diff.item():.8f} " + f"grad_norm_reference={reference_grad_norm.item():.8f} " + f"grad_norm_candidate={candidate_grad_norm.item():.8f} " + f"grad_norm_diff={grad_norm_diff.item():.8f}", + flush=True, + ) + + torch.testing.assert_close( + candidate_avg.float(), + reference_avg.float(), + atol=_LM_LOSS_ATOL, + rtol=0.0, + ) + torch.testing.assert_close( + candidate_mtp_losses, + reference_mtp_losses, + atol=_MTP_LOSS_ATOL, + rtol=0.0, + ) + torch.testing.assert_close( + candidate_grad_norm, + reference_grad_norm, + atol=_GRAD_NORM_ATOL, + rtol=_GRAD_NORM_RTOL, + ) + finally: + _destroy_model_parallel_without_barrier() diff --git a/tests/unit_tests/ssm/test_hybrid_block.py b/tests/unit_tests/ssm/test_hybrid_block.py index 48946d0b42f..f098962d3ce 100644 --- a/tests/unit_tests/ssm/test_hybrid_block.py +++ b/tests/unit_tests/ssm/test_hybrid_block.py @@ -302,6 +302,51 @@ def test_hyper_connection_gpu_forward(self): assert output.shape[2] == block.config.hidden_size assert output.dtype == torch.float32 + def test_hyper_connection_mlp_fast_path_forwards_packed_seq_params(self, monkeypatch): + """mHC MLP-only fast path preserves packed sequence metadata for MoE routing.""" + block = self.get_mamba_block(Symbols.MLP, enable_hyper_connections=True) + layer = block.layers[0] + inner = layer.inner_layer + + captured = {} + + def fake_forward_mlp_output_with_bias( + hidden_states, + inference_context=None, + padding_mask=None, + input_ids=None, + packed_seq_params=None, + ): + captured["packed_seq_params"] = packed_seq_params + captured["padding_mask"] = padding_mask + captured["input_ids"] = input_ids + return (hidden_states, None), hidden_states + + monkeypatch.setattr( + inner, "_forward_mlp_output_with_bias", fake_forward_mlp_output_with_bias + ) + + hidden_states = torch.ones((4, 1, block.config.hidden_size)) + padding_mask = torch.zeros((1, 4), dtype=torch.bool) + input_ids = torch.arange(4).view(1, 4) + packed_seq_params = object() + + result = layer._call_inner_transformer_layer_without_local_bda( + hidden_states=hidden_states, + attention_mask=None, + inference_context=None, + rotary_pos_emb=None, + sequence_len_offset=None, + packed_seq_params=packed_seq_params, + padding_mask=padding_mask, + input_ids=input_ids, + ) + + assert result is not None + assert captured["packed_seq_params"] is packed_seq_params + assert captured["padding_mask"] is padding_mask + assert captured["input_ids"] is input_ids + def test_hyper_connection_gdn_gpu_forward(self): """mHC runs through GDN, attention, and Mamba hybrid layers.""" layer_pattern = Symbols.GDN + Symbols.ATTENTION + Symbols.MAMBA diff --git a/tests/unit_tests/test_context_parallel_layout.py b/tests/unit_tests/test_context_parallel_layout.py index b762e4594a1..749c5a56974 100644 --- a/tests/unit_tests/test_context_parallel_layout.py +++ b/tests/unit_tests/test_context_parallel_layout.py @@ -1,15 +1,167 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +from types import SimpleNamespace + import pytest import torch -from megatron.core.context_parallel_layout import get_thd_context_parallel_rank_indices +import megatron.core.context_parallel_layout as context_parallel_layout +import megatron.core.context_parallel_layout.conversion as context_parallel_layout_conversion +from megatron.core import parallel_state +from megatron.core.context_parallel_layout import ( + CpPartitionModeConverter, + build_thd_cp_partition_route, + decode_thd_cp_partition_route, + get_context_parallel_layout_chunk_indices, + get_preferred_cp_partition_mode_for_layer, + get_stage_entry_partition_mode, + get_thd_context_parallel_rank_indices, + get_thd_cp_partition_route, + prebuild_thd_cp_partition_routes, + replace_packed_seq_params_cp_partition_mode, +) +from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_stage_input_cp_partition_mode, +) +from megatron.core.models.hybrid.hybrid_layer_allocation import ( + get_hybrid_stage_input_cp_partition_mode, +) +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.transformer.transformer_config import TransformerConfig +from tests.unit_tests.test_utilities import Utils + + +class _PipelineLayout: + + def __init__(self, offset): + self.offset = offset + + def get_layer_offset(self, **_kwargs): + return self.offset + + +def _minimal_transformer_config_kwargs(): + return dict(num_layers=1, hidden_size=8, num_attention_heads=1) + + +def test_transformer_config_cp_partition_mode_selector(): + assert TransformerConfig(**_minimal_transformer_config_kwargs()).cp_partition_mode == "zigzag" + assert ( + TransformerConfig( + **_minimal_transformer_config_kwargs(), cp_partition_mode="contiguous" + ).cp_partition_mode + == "contiguous" + ) + assert ( + TransformerConfig(**_minimal_transformer_config_kwargs(), cp_partition_mode="auto") + .cp_partition_mode + == "auto" + ) + + with pytest.raises(ValueError, match="Unsupported cp_partition_mode"): + TransformerConfig(**_minimal_transformer_config_kwargs(), cp_partition_mode="invalid") + + +def test_packed_seq_params_rejects_auto_cp_partition_mode(): + with pytest.raises(ValueError, match="concrete runtime layout"): + PackedSeqParams(qkv_format="sbhd", cp_partition_mode="auto") + + +class IdentityOp: + + def get_preferred_cp_partition_mode(self): + return None + + +class GatedDeltaNet: + + def get_preferred_cp_partition_mode(self): + mode = getattr(self.config, "linear_cp_mode", "chunkwise") + if mode == "chunkwise": + return "contiguous" + if mode == "headwise": + return "zigzag" + raise ValueError(f"Unsupported GatedDeltaNet linear_cp_mode: {mode!r}.") + + +class _PreferredContiguousLayer: + + def get_preferred_cp_partition_mode(self): + return "contiguous" + + +class _PreferredAgnosticLayer: + + def get_preferred_cp_partition_mode(self): + return None + + +class _DynamicPreferredLayer: + + def get_preferred_cp_partition_mode(self): + return "zigzag" + + +class _InvalidPreferredLayer: + + def get_preferred_cp_partition_mode(self): + return "unsupported" + + +class _FakeGroup: + + def __init__(self, size, rank): + self._size = size + self._rank = rank + + def size(self): + return self._size + + def rank(self): + return self._rank def _token_ranges(*spans): return [token for start, end in spans for token in range(start, end)] +def _make_sequence_tensor(total_seq_len, seq_dim, device): + if seq_dim == 0: + shape = (total_seq_len, 3, 5) + elif seq_dim == 1: + shape = (3, total_seq_len, 5) + else: + raise ValueError(f"Unsupported test seq_dim {seq_dim}.") + return torch.arange(torch.prod(torch.tensor(shape)), device=device, dtype=torch.float32).view( + *shape + ) + + +def _get_sequence_parallel_shard(tensor, seq_dim, tp_group): + tp_size = tp_group.size() + tp_rank = tp_group.rank() + assert tensor.size(seq_dim) % tp_size == 0 + return tensor.chunk(tp_size, dim=seq_dim)[tp_rank].contiguous() + + +def _get_sbhd_tensor_on_this_cp_rank(tensor, seq_dim, cp_group, cp_partition_mode): + cp_size = cp_group.size() + cp_rank = cp_group.rank() + cp_idx = get_context_parallel_layout_chunk_indices( + cp_size, cp_rank, cp_partition_mode + ).to(device=tensor.device) + tensor = tensor.view( + *tensor.shape[:seq_dim], 2 * cp_size, -1, *tensor.shape[(seq_dim + 1) :] + ) + tensor = tensor.index_select(seq_dim, cp_idx) + return tensor.view(*tensor.shape[:seq_dim], -1, *tensor.shape[(seq_dim + 2) :]) + + +def test_context_parallel_layout_chunk_indices(): + assert get_context_parallel_layout_chunk_indices(4, 2, "zigzag").tolist() == [2, 5] + assert get_context_parallel_layout_chunk_indices(4, 2, "contiguous").tolist() == [4, 5] + + def test_thd_context_parallel_rank_indices_match_per_sequence_chunk_order(): cu_seqlens = torch.tensor([0, 16, 40]) @@ -58,6 +210,248 @@ def test_thd_context_parallel_rank_indices_reject_uneven_chunks(): get_thd_context_parallel_rank_indices(torch.tensor([0, 10]), 2, 0, "zigzag") +def test_thd_contiguous_rank_indices_allow_uneven_sequence_lengths(): + cu_seqlens = torch.tensor([0, 10, 18]) + + assert get_thd_context_parallel_rank_indices(cu_seqlens, 2, 0, "contiguous").tolist() == list( + range(0, 9) + ) + assert get_thd_context_parallel_rank_indices(cu_seqlens, 2, 1, "contiguous").tolist() == list( + range(9, 18) + ) + + +@pytest.mark.parametrize( + ("source_layout", "target_layout"), [("zigzag", "contiguous"), ("contiguous", "zigzag")] +) +@pytest.mark.parametrize( + ("cu_seqlens", "cp_size"), + [ + (torch.tensor([0, 16, 40]), 2), + (torch.tensor([0, 32, 96, 128]), 4), + (torch.tensor([0, 32, 96, 128, 128, 128]), 4), + ], +) +def test_thd_cp_partition_route_reassembles_target_layout( + source_layout, target_layout, cu_seqlens, cp_size +): + source_indices = [ + get_thd_context_parallel_rank_indices(cu_seqlens, cp_size, rank, source_layout) + for rank in range(cp_size) + ] + target_indices = [ + get_thd_context_parallel_rank_indices(cu_seqlens, cp_size, rank, target_layout) + for rank in range(cp_size) + ] + routes = [ + build_thd_cp_partition_route(cu_seqlens, cp_size, rank, source_layout, target_layout) + for rank in range(cp_size) + ] + decoded_routes = [ + decode_thd_cp_partition_route(routes[rank], cp_size, rank) for rank in range(cp_size) + ] + for rank, ( + local_source_length, + local_target_length, + _, + _, + _, + _, + ) in enumerate(decoded_routes): + assert local_source_length == source_indices[rank].numel() + assert local_target_length == target_indices[rank].numel() + send_buffers = [] + for rank, (_, _, send_rows, _, _, _) in enumerate(decoded_routes): + send_buffers.append( + source_indices[rank] + if send_rows is None + else source_indices[rank].index_select(0, send_rows) + ) + + for dst_rank in range(cp_size): + recv_chunks = [] + for src_rank in range(cp_size): + _, _, _, _, input_split_sizes, _ = decoded_routes[src_rank] + send_offset = sum(input_split_sizes[:dst_rank]) + send_len = input_split_sizes[dst_rank] + recv_chunks.append(send_buffers[src_rank].narrow(0, send_offset, send_len)) + recv_buf = torch.cat(recv_chunks, dim=0) + _, local_target_length, _, recv_rows, _, _ = decoded_routes[dst_rank] + if recv_rows is None: + out = recv_buf + else: + out = torch.empty(local_target_length, dtype=recv_buf.dtype) + out.index_copy_(0, recv_rows, recv_buf) + assert torch.equal(out, target_indices[dst_rank]) + + +@pytest.mark.internal +@pytest.mark.parametrize( + ("source_layout", "target_layout"), [("zigzag", "contiguous"), ("contiguous", "zigzag")] +) +@pytest.mark.parametrize("seq_dim", [0, 1]) +def test_sbhd_convert_cp_partition_mode_matches_direct_target_shard( + source_layout, target_layout, seq_dim +): + if not torch.cuda.is_available() or Utils.world_size < 2: + pytest.skip("SBHD CP partition-mode conversion needs at least two CUDA ranks.") + + cp_size = 2 + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=cp_size) + try: + cp_group = parallel_state.get_context_parallel_group() + full_tensor = _make_sequence_tensor( + total_seq_len=32, + seq_dim=seq_dim, + device=torch.device(f"cuda:{torch.cuda.current_device()}"), + ) + source_shard = _get_sbhd_tensor_on_this_cp_rank( + full_tensor, seq_dim, cp_group, cp_partition_mode=source_layout + ) + + converted = context_parallel_layout.convert_cp_partition_mode( + source_shard, + cp_group, + source_partition_mode=source_layout, + target_partition_mode=target_layout, + seq_dim=seq_dim, + ) + expected = _get_sbhd_tensor_on_this_cp_rank( + full_tensor, seq_dim, cp_group, cp_partition_mode=target_layout + ) + + torch.testing.assert_close(converted, expected, atol=0.0, rtol=0.0) + finally: + Utils.destroy_model_parallel() + + +@pytest.mark.internal +@pytest.mark.parametrize( + ("source_layout", "target_layout", "seq_dim", "sequence_parallel"), + [ + pytest.param("zigzag", "contiguous", 0, False, id="zigzag-contiguous-seq0"), + pytest.param("zigzag", "contiguous", 1, False, id="zigzag-contiguous-seq1"), + pytest.param("contiguous", "zigzag", 0, False, id="contiguous-zigzag-seq0"), + pytest.param("contiguous", "zigzag", 1, False, id="contiguous-zigzag-seq1"), + pytest.param("zigzag", "contiguous", 0, True, id="sequence-parallel"), + ], +) +def test_sbhd_convert_cp_partition_mode_backward_matches_direct_source_shard( + source_layout, target_layout, seq_dim, sequence_parallel +): + min_world_size = 4 if sequence_parallel else 2 + if not torch.cuda.is_available() or Utils.world_size < min_world_size: + pytest.skip( + f"SBHD CP partition-mode conversion backward needs at least {min_world_size} " + "CUDA ranks." + ) + + cp_size = 2 + tp_size = 2 if sequence_parallel else 1 + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp_size, context_parallel_size=cp_size + ) + try: + cp_group = parallel_state.get_context_parallel_group() + tp_group = parallel_state.get_tensor_model_parallel_group() if sequence_parallel else None + full_tensor = _make_sequence_tensor( + total_seq_len=32, + seq_dim=seq_dim, + device=torch.device(f"cuda:{torch.cuda.current_device()}"), + ) + full_upstream_grad = full_tensor.mul(0.125).add(1.0) + source_shard = _get_sbhd_tensor_on_this_cp_rank( + full_tensor, seq_dim, cp_group, cp_partition_mode=source_layout + ) + if sequence_parallel: + source_shard = _get_sequence_parallel_shard(source_shard, seq_dim, tp_group) + source_shard = source_shard.detach().requires_grad_(True) + + convert_kwargs = ( + {"sequence_parallel": True, "tp_group": tp_group} if sequence_parallel else {} + ) + converted = context_parallel_layout.convert_cp_partition_mode( + source_shard, + cp_group, + source_partition_mode=source_layout, + target_partition_mode=target_layout, + seq_dim=seq_dim, + **convert_kwargs, + ) + target_upstream_grad = _get_sbhd_tensor_on_this_cp_rank( + full_upstream_grad, seq_dim, cp_group, cp_partition_mode=target_layout + ) + if sequence_parallel: + target_upstream_grad = _get_sequence_parallel_shard( + target_upstream_grad, seq_dim, tp_group + ) + converted.mul(target_upstream_grad).sum().backward() + expected_source_grad = _get_sbhd_tensor_on_this_cp_rank( + full_upstream_grad, seq_dim, cp_group, cp_partition_mode=source_layout + ) + if sequence_parallel: + expected_source_grad = _get_sequence_parallel_shard( + expected_source_grad, seq_dim, tp_group + ) + + torch.testing.assert_close(source_shard.grad, expected_source_grad, atol=0.0, rtol=0.0) + finally: + Utils.destroy_model_parallel() + + +def test_prebuild_thd_cp_partition_routes_populates_direct_fields(): + packed_seq_params = SimpleNamespace( + qkv_format="thd", + cu_seqlens_q=torch.tensor([0, 16, 40]), + cu_seqlens_q_padded=None, + cp_partition_route_zigzag_to_contiguous=None, + cp_partition_route_contiguous_to_zigzag=None, + ) + cp_group = _FakeGroup(size=2, rank=0) + prebuild_thd_cp_partition_routes(packed_seq_params, cp_group) + + route = get_thd_cp_partition_route(packed_seq_params, "zigzag", "contiguous") + same_route = get_thd_cp_partition_route(packed_seq_params, "zigzag", "contiguous") + reverse_route = get_thd_cp_partition_route(packed_seq_params, "contiguous", "zigzag") + + assert same_route is route + assert reverse_route is not route + assert packed_seq_params.cp_partition_route_zigzag_to_contiguous is route + assert packed_seq_params.cp_partition_route_contiguous_to_zigzag is reverse_route + + +def test_prebuild_thd_cp_partition_routes_is_best_effort(): + packed_seq_params = SimpleNamespace( + qkv_format="thd", + cu_seqlens_q=torch.tensor([0, 10, 18]), + cu_seqlens_q_padded=None, + cp_partition_route_zigzag_to_contiguous=None, + cp_partition_route_contiguous_to_zigzag=None, + ) + cp_group = _FakeGroup(size=2, rank=0) + + prebuild_thd_cp_partition_routes(packed_seq_params, cp_group) + + assert packed_seq_params.cp_partition_route_zigzag_to_contiguous is None + assert packed_seq_params.cp_partition_route_contiguous_to_zigzag is None + + +def test_cp_partition_mode_annotation_preserves_route_tensor_identity(): + route = torch.tensor([2, 0, 4, 4, 0, 0, 2, 2, 2, 2]) + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cp_partition_mode="zigzag", + cp_partition_route_zigzag_to_contiguous=route, + ) + + updated = replace_packed_seq_params_cp_partition_mode(packed_seq_params, "contiguous") + + assert updated is packed_seq_params + assert updated.cp_partition_mode == "contiguous" + assert updated.cp_partition_route_zigzag_to_contiguous is route + assert replace_packed_seq_params_cp_partition_mode(updated, "contiguous") is updated + + def test_thd_context_parallel_rank_indices_reject_decreasing_boundaries(): with pytest.raises(ValueError, match="nondecreasing"): get_thd_context_parallel_rank_indices(torch.tensor([0, 16, 8]), 2, 0, "zigzag") @@ -66,3 +460,200 @@ def test_thd_context_parallel_rank_indices_reject_decreasing_boundaries(): def test_thd_context_parallel_rank_indices_reject_unknown_layout(): with pytest.raises(ValueError, match="Unsupported"): get_thd_context_parallel_rank_indices(torch.tensor([0, 16]), 2, 0, "interleaved") + + +def test_cp_partition_mode_converter_recurses_over_tensor_containers(monkeypatch): + calls = [] + + def fake_convert(tensor, cp_group, **kwargs): + calls.append((tensor, cp_group, kwargs)) + return tensor + 10 + + monkeypatch.setattr( + context_parallel_layout_conversion, "convert_cp_partition_mode", fake_convert + ) + cp_group = SimpleNamespace(size=lambda: 2) + config = SimpleNamespace(cuda_graph_impl=None) + cu_seqlens = torch.tensor([0, 8]) + untouched = object() + value = (torch.tensor([1]), [None, untouched, torch.tensor([2])]) + + converter = CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=SimpleNamespace( + qkv_format="thd", + cu_seqlens_q=cu_seqlens, + cu_seqlens_q_padded=None, + cp_partition_route_zigzag_to_contiguous=torch.tensor([0]), + ), + source_partition_mode="zigzag", + target_partition_mode="contiguous", + config=config, + ) + converted = converter.convert(value, seq_dim=lambda tensor: tensor.dim() - 1) + + assert torch.equal(converted[0], torch.tensor([11])) + assert converted[1][0] is None + assert converted[1][1] is untouched + assert torch.equal(converted[1][2], torch.tensor([12])) + assert [call[1] for call in calls] == [cp_group, cp_group] + assert [call[2]["seq_dim"] for call in calls] == [0, 0] + assert all(call[2]["cu_seqlens"] is cu_seqlens for call in calls) + + +def test_cp_partition_mode_converter_rejects_thd_full_iteration_cuda_graph_conversion(): + cp_group = SimpleNamespace(size=lambda: 2) + packed_seq_params = SimpleNamespace(qkv_format="thd") + config = SimpleNamespace(cuda_graph_impl="full_iteration") + + CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode="zigzag", + target_partition_mode="zigzag", + config=config, + ) + + with pytest.raises(ValueError, match="Full-iteration CUDA graph"): + CpPartitionModeConverter( + cp_group=cp_group, + packed_seq_params=packed_seq_params, + source_partition_mode="zigzag", + target_partition_mode="contiguous", + config=config, + ) + + +def test_preferred_partition_mode_rejects_unknown_layer_type(): + with pytest.raises(ValueError, match="Cannot determine CP partition mode"): + get_preferred_cp_partition_mode_for_layer(object(), SimpleNamespace(cp_comm_type=None)) + + +def test_preferred_partition_mode_uses_explicit_method(): + config = SimpleNamespace(cp_comm_type=None) + + assert ( + get_preferred_cp_partition_mode_for_layer(_PreferredContiguousLayer(), config) + == "contiguous" + ) + assert get_preferred_cp_partition_mode_for_layer(_PreferredAgnosticLayer(), config) is None + + +def test_preferred_partition_mode_uses_dynamic_method(): + config = SimpleNamespace(cp_comm_type=None) + + assert get_preferred_cp_partition_mode_for_layer(_DynamicPreferredLayer(), config) == "zigzag" + + +def test_preferred_partition_mode_rejects_invalid_explicit_declaration(): + config = SimpleNamespace(cp_comm_type=None) + + with pytest.raises(ValueError, match="Invalid CP partition mode preference"): + get_preferred_cp_partition_mode_for_layer(_InvalidPreferredLayer(), config) + + +@pytest.mark.parametrize( + ("linear_cp_mode", "expected_partition_mode"), + [("chunkwise", "contiguous"), ("headwise", "zigzag")], +) +def test_gated_delta_net_layer_layout_policy_follows_cp_mode( + linear_cp_mode, expected_partition_mode +): + config = SimpleNamespace(cp_comm_type=None, linear_cp_mode=linear_cp_mode) + layer = object.__new__(GatedDeltaNet) + layer.config = config + + assert get_preferred_cp_partition_mode_for_layer(layer, config) == expected_partition_mode + + +def test_get_stage_entry_partition_mode_uses_packed_metadata(): + packed_seq_params = SimpleNamespace( + cp_partition_mode="zigzag", + cp_group=_FakeGroup(size=2, rank=0), + ) + + assert ( + get_stage_entry_partition_mode( + packed_seq_params, "zigzag", owner_name="TestBlock" + ) + == "zigzag" + ) + + +def test_get_stage_entry_partition_mode_rejects_mismatch(): + packed_seq_params = SimpleNamespace( + cp_partition_mode="contiguous", + cp_group=_FakeGroup(size=2, rank=0), + ) + + with pytest.raises(AssertionError, match="expected CP stage entry partition mode"): + get_stage_entry_partition_mode( + packed_seq_params, "zigzag", owner_name="TestBlock" + ) + + +def test_get_stage_entry_partition_mode_requires_layout_under_cp(): + packed_seq_params = SimpleNamespace(cp_partition_mode=None, cp_group=_FakeGroup(size=2, rank=0)) + + with pytest.raises(ValueError, match="requires a CP stage entry partition mode"): + get_stage_entry_partition_mode(packed_seq_params, None, owner_name="TestBlock") + + +def test_get_stage_entry_partition_mode_allows_missing_layout_without_cp(): + packed_seq_params = SimpleNamespace(cp_partition_mode=None, cp_group=_FakeGroup(size=1, rank=0)) + + assert get_stage_entry_partition_mode(packed_seq_params, None, owner_name="TestBlock") is None + + +def test_gated_delta_net_chunkwise_layout_plan_follows_linear_attention_pattern(): + config = SimpleNamespace( + experimental_attention_variant="gated_delta_net", + linear_attention_freq=2, + linear_cp_mode="chunkwise", + num_layers=4, + pipeline_model_parallel_layout=None, + pipeline_model_parallel_size=1, + ) + + assert get_experimental_attention_variant_stage_input_cp_partition_mode(config) == "contiguous" + + config.pipeline_model_parallel_layout = _PipelineLayout(offset=2) + assert get_experimental_attention_variant_stage_input_cp_partition_mode(config) == "zigzag" + + +def test_gated_delta_net_headwise_layout_plan_uses_zigzag(): + config = SimpleNamespace( + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1, 0], + linear_cp_mode="headwise", + num_layers=2, + pipeline_model_parallel_layout=None, + pipeline_model_parallel_size=1, + ) + + assert get_experimental_attention_variant_stage_input_cp_partition_mode(config) == "zigzag" + + +@pytest.mark.parametrize( + ("layer_pattern", "local_index", "linear_cp_mode", "attention_variant", "expected"), + [ + pytest.param("M-G", 0, "chunkwise", None, "zigzag", id="before-previous-mamba"), + pytest.param("M-G", 2, "chunkwise", None, "zigzag", id="before-previous-gdn"), + pytest.param("M-G", 3, "chunkwise", None, "contiguous", id="after-chunkwise-gdn"), + pytest.param("-G", 0, "chunkwise", None, "contiguous", id="future-chunkwise-gdn"), + pytest.param("-G", 0, "headwise", None, "zigzag", id="future-headwise-gdn"), + pytest.param("D-E", 0, "chunkwise", "dsv4_hybrid", "contiguous", id="dsv4-d"), + pytest.param("C-E", 0, "chunkwise", "dsv4_hybrid", "contiguous", id="dsv4-c"), + pytest.param("-E-", 0, "chunkwise", None, None, id="no-sensitive-layer-start"), + pytest.param("-E-", 3, "chunkwise", None, None, id="no-sensitive-layer-late"), + ], +) +def test_hybrid_stage_input_layout_policy( + layer_pattern, local_index, linear_cp_mode, attention_variant, expected +): + config = SimpleNamespace( + experimental_attention_variant=attention_variant, + linear_cp_mode=linear_cp_mode, + ) + + assert get_hybrid_stage_input_cp_partition_mode(config, layer_pattern, local_index) == expected diff --git a/tests/unit_tests/test_sequence_packing.py b/tests/unit_tests/test_sequence_packing.py index f0467ca8bcd..e7299006870 100644 --- a/tests/unit_tests/test_sequence_packing.py +++ b/tests/unit_tests/test_sequence_packing.py @@ -266,7 +266,6 @@ def test_dsv4_thd_dynamic_cp_pads_before_slicing( cp=_MockCPGroup(size=4, rank=cp_rank), ) config = SimpleNamespace( - cp_partition_mode="contiguous", pad_packed_seq_alignment=alignment, max_seqlen_per_dp_cp_rank=8, thd_max_packed_sequences=None, @@ -338,7 +337,11 @@ def record_slice(batch, cp_group, **kwargs): ) result = get_batch_on_this_rank_for_sequence_packing( - data_iterator=iter([batch]), dynamic_cp=True, pg_collection=pg_collection, config=config + data_iterator=iter([batch]), + dynamic_cp=True, + pg_collection=pg_collection, + config=config, + cp_partition_mode="contiguous", ) local_tokens, _, _, _, _, packed_seq_params, padding_mask = result @@ -402,7 +405,6 @@ def test_non_dummy_zigzag_padding_updates_metadata_before_cp_slice(monkeypatch): cp=_MockCPGroup(size=cp_size, rank=cp_rank), ) config = SimpleNamespace( - cp_partition_mode="zigzag", pad_packed_seq_alignment=4, max_seqlen_per_dp_cp_rank=local_target, thd_max_packed_sequences=None, @@ -447,7 +449,11 @@ def get_padded_zigzag_indices(cu_seqlens, total, world_size, rank): ) result = get_batch_on_this_rank_for_sequence_packing( - data_iterator=iter([batch]), dynamic_cp=True, pg_collection=pg_collection, config=config + data_iterator=iter([batch]), + dynamic_cp=True, + pg_collection=pg_collection, + config=config, + cp_partition_mode="zigzag", ) local_tokens, _, _, _, _, packed_seq_params, padding_mask = result @@ -592,12 +598,16 @@ def test_get_batch_on_this_rank_for_sequence_packing(tp, pp, cp, dynamic_cp, loc else: data_iterator = None + effective_cp = local_cp_size if dynamic_cp else cp + cp_partition_mode = "zigzag" if effective_cp is not None and effective_cp > 1 else None + # Call the function under test result = get_batch_on_this_rank_for_sequence_packing( data_iterator=data_iterator, mtp_on_this_rank=False, vp_stage=None, dynamic_cp=dynamic_cp, + cp_partition_mode=cp_partition_mode, ) # The helper returns a 7-tuple; scheduler THD always provides padding_mask. @@ -714,7 +724,6 @@ def test_get_batch_on_this_rank_for_sequence_packing(tp, pp, cp, dynamic_cp, loc # ===================================================================== # TEST 4: Verify CP partitioning # ===================================================================== - effective_cp = local_cp_size if dynamic_cp else cp if effective_cp is not None and effective_cp > 1: expected_seq_len = args.seq_length // effective_cp diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_dsv4_hybrid_attention_cp.py b/tests/unit_tests/transformer/experimental_attention_variant/test_dsv4_hybrid_attention_cp.py index 09e40725bf7..e928695e6bf 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_dsv4_hybrid_attention_cp.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_dsv4_hybrid_attention_cp.py @@ -26,6 +26,7 @@ _SEED, _build_attention, _make_config, + patch_hadamard_if_needed, # noqa: F401 ) from tests.unit_tests.transformer.experimental_attention_variant.test_dsv4_hybrid_native_parity import ( _DSV4_VARIANTS, @@ -293,7 +294,6 @@ def _make_dsv4_cp_config( dsa_indexer_loss_coeff=dsa_indexer_loss_coeff, dsa_indexer_use_sparse_loss=dsa_indexer_use_sparse_loss, context_parallel_size=context_parallel_size, - cp_partition_mode="contiguous" if context_parallel_size > 1 else "zigzag", sequence_packing_scheduler="dp_balanced" if context_parallel_size > 1 else None, csa_dense_mode=False, csa_compress_rotary_base=shape["csa_compress_rotary_base"], diff --git a/tests/unit_tests/transformer/test_attention.py b/tests/unit_tests/transformer/test_attention.py index 6a65a91a2fc..6893da6c3a8 100644 --- a/tests/unit_tests/transformer/test_attention.py +++ b/tests/unit_tests/transformer/test_attention.py @@ -10,10 +10,11 @@ from torch.nn import functional as F import megatron.core.parallel_state as parallel_state -from megatron.core.hyper_comm_grid import HyperCommGrid -from megatron.core.models.common.embeddings.rope_utils import ( - get_pos_emb_on_this_cp_rank as get_tensor_on_this_cp_rank, +from megatron.core.context_parallel_layout import ( + get_context_parallel_layout_chunk_indices, + get_thd_context_parallel_rank_indices, ) +from megatron.core.hyper_comm_grid import HyperCommGrid from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_local_spec, get_gpt_layer_with_transformer_engine_spec, @@ -715,6 +716,9 @@ def _test_parallel_attention_correctness( sequence_length=256, micro_batch_size=4, sequence_packing=False, + cp_partition_mode="zigzag", + cp_stage_entry_partition_mode="zigzag", + compare_param_grads=False, ): # Model initialization function def initialize_gpt_model( @@ -729,6 +733,7 @@ def initialize_gpt_model( post_process=post_process, vp_stage=vp_stage, pg_collection=pg_collection, + cp_stage_entry_partition_mode=cp_stage_entry_partition_mode, ) return gpt_model @@ -766,8 +771,24 @@ def initialize_gpt_model( mock_args.no_load_rng = True save_checkpoint(10, gpt_model, None, None, 0) + def get_param_grad(param): + grad = param.grad + if grad is None: + grad = getattr(param, "main_grad", None) + if grad is None: + return torch.zeros_like(param, dtype=torch.float32) + return grad + + def zero_param_grads(module): + for param in module.parameters(): + param.grad = None + main_grad = getattr(param, "main_grad", None) + if main_grad is not None: + main_grad.zero_() + # Calculate baseline output attention = gpt_model[0].decoder.layers[0].self_attention + zero_param_grads(attention) output_hidden_states_baseline, bias_hidden_states_baseline = attention( input_hidden_states, attention_mask=None ) @@ -775,6 +796,13 @@ def initialize_gpt_model( # Save baseline output input_grad_baseline = input_hidden_states.grad.detach() + param_grads_baseline = None + if compare_param_grads: + param_grads_baseline = { + name: get_param_grad(param).detach().float().clone() + for name, param in attention.named_parameters() + if param.requires_grad + } output_hidden_states_baseline = output_hidden_states_baseline.detach() bias_hidden_states_baseline = bias_hidden_states_baseline if bias_hidden_states_baseline is not None: @@ -806,10 +834,29 @@ def initialize_gpt_model( tp_rank = parallel_state.get_tensor_model_parallel_rank() def get_tensor_on_this_rank(tensor): - if cp > 1: - tensor = get_tensor_on_this_cp_rank(tensor, 0, cp_group) if sequence_packing: tensor = tensor.transpose(0, 1).contiguous().view(-1, 1, *tensor.shape[2:]) + if cp > 1: + cu_seqlens_tensor = torch.tensor( + [i * sequence_length for i in range(micro_batch_size + 1)], + device=tensor.device, + ) + cp_rank = torch.distributed.get_rank(cp_group) + index = get_thd_context_parallel_rank_indices( + cu_seqlens_tensor, cp, cp_rank, cp_partition_mode + ) + tensor = tensor.index_select(0, index) + elif cp > 1: + cp_rank = torch.distributed.get_rank(cp_group) + cp_idx = get_context_parallel_layout_chunk_indices( + cp, cp_rank, cp_partition_mode + ).to(device=tensor.device) + tensor = tensor.view( + 2 * cp, + -1, + *tensor.shape[1:], + ) + tensor = tensor.index_select(0, cp_idx).view(-1, *tensor.shape[2:]) if tp > 1 and sp: sp_seg = tensor.shape[0] // tp tensor = tensor[tp_rank * sp_seg : (tp_rank + 1) * sp_seg] @@ -819,15 +866,27 @@ def get_tensor_on_this_rank(tensor): if sequence_packing: cu_seqlens = [i * sequence_length for i in range(micro_batch_size + 1)] packed_seq_params = make_test_packed_seq_params(cu_seqlens=cu_seqlens) + packed_seq_params.cp_partition_mode = cp_partition_mode else: packed_seq_params = None input_hidden_states = get_tensor_on_this_rank(input_hidden_states) input_hidden_states = input_hidden_states.detach().requires_grad_(True) parallel_attention = gpt_model[0].decoder.layers[0].self_attention + zero_param_grads(parallel_attention) output_hidden_states_parallel, bias_hidden_states_parallel = parallel_attention( input_hidden_states, attention_mask=None, packed_seq_params=packed_seq_params ) output_hidden_states_parallel.sum().backward() + param_grads_parallel = None + if compare_param_grads: + assert tp == 1, "Parameter gradient parity only supports unsharded parameters." + param_grads_parallel = {} + for name, param in parallel_attention.named_parameters(): + if param.requires_grad: + grad = get_param_grad(param) + if cp > 1: + torch.distributed.all_reduce(grad, group=cp_group) + param_grads_parallel[name] = grad.detach().float().clone() input_grad_parallel = input_hidden_states.grad.detach() # Check if the output is close @@ -904,6 +963,12 @@ def assert_close_or_cosine_similarity(baseline, parallel, tensor_name): assert_close_or_cosine_similarity( bias_hidden_states_baseline, bias_hidden_states_parallel, "bias_hidden_states" ) + if compare_param_grads: + assert param_grads_baseline.keys() == param_grads_parallel.keys() + for name, grad_baseline in param_grads_baseline.items(): + assert_close_or_cosine_similarity( + grad_baseline, param_grads_parallel[name], f"param_grad[{name}]" + ) Utils.destroy_model_parallel() diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index 390b5a6109f..2a7afcd7b3c 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -55,6 +55,10 @@ _SEED = 42 +def _make_sbhd_cp_packed_seq_params(cp_group): + return PackedSeqParams(qkv_format="sbhd", cp_group=cp_group, cp_partition_mode="zigzag") + + class TestMultiTokenPredictionLayer: def setup_method(self, method): os.environ['CUDA_DEVICE_MAX_CONNECTIONS'] = '1' @@ -739,8 +743,12 @@ def set_ckpt_path(ckpt_path): load_checkpoint(gpt_model, optimizer, opt_param_scheduler, strict=False) batch["output_ref"] = output_ref # Get batch for current CP rank (handles CP tensor splitting) + cp_group = get_context_parallel_group() batch = get_batch_on_this_cp_rank( - batch, is_hybrid_cp=False, cp_group=get_context_parallel_group() + batch, + is_hybrid_cp=False, + cp_group=cp_group, + cp_partition_mode="zigzag", ) tokens, labels, loss_mask, attention_mask, position_ids, output_ref = batch.values() output = gpt_model[0].forward( @@ -749,6 +757,9 @@ def set_ckpt_path(ckpt_path): attention_mask=attention_mask, labels=labels, loss_mask=loss_mask, + packed_seq_params=( + _make_sbhd_cp_packed_seq_params(cp_group) if cp > 1 else None + ), ) # Combine normalized loss contributions across DP+CP. MTPLossLoggingHelper.reduce_loss_in_tracker() @@ -1098,6 +1109,7 @@ def test_roll_tensor_with_packed_sequences(self, cp): max_seqlen_q=6, # max(4, 6) - max local seq length per sequence max_seqlen_kv=6, qkv_format='thd', + cp_partition_mode="zigzag", ) # Roll by -1 (shift left) with CP communication @@ -1270,6 +1282,7 @@ def test_roll_tensor_with_packed_sequences_odd_seqlen(self, cp): max_seqlen_q=11, max_seqlen_kv=11, qkv_format='thd', + cp_partition_mode="zigzag", ) rolled, sum_val = roll_tensor( @@ -1798,7 +1811,10 @@ def set_ckpt_path(ckpt_path): batch["output_ref"] = output_ref batch = get_batch_on_this_cp_rank( - batch, is_hybrid_cp=False, cp_group=get_context_parallel_group() + batch, + is_hybrid_cp=False, + cp_group=get_context_parallel_group(), + cp_partition_mode="zigzag", ) tokens, labels, loss_mask, attention_mask, position_ids, output_ref = batch.values() output = mamba_model[0].forward( diff --git a/tests/unit_tests/transformer/test_thd_correctness.py b/tests/unit_tests/transformer/test_thd_correctness.py index ac759c2bdce..f4cbe491619 100644 --- a/tests/unit_tests/transformer/test_thd_correctness.py +++ b/tests/unit_tests/transformer/test_thd_correctness.py @@ -217,6 +217,7 @@ def to_cu_seqlens(lens): max_seqlen_q=max(padded), max_seqlen_kv=max(padded), qkv_format='thd', + cp_partition_mode="zigzag" if cp_size > 1 else None, ) diff --git a/tests/unit_tests/transformer/test_thd_cuda_graph.py b/tests/unit_tests/transformer/test_thd_cuda_graph.py index a3abe63a2f2..8ab9344935b 100644 --- a/tests/unit_tests/transformer/test_thd_cuda_graph.py +++ b/tests/unit_tests/transformer/test_thd_cuda_graph.py @@ -423,9 +423,11 @@ def test_cp_dummy_rejects_tail_that_cannot_be_partitioned(self): Utils.destroy_model_parallel() Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=2) + psp = _make_psp([4]) + psp.cp_partition_mode = "zigzag" with pytest.raises(AssertionError, match="must be divisible"): pad_sequence_for_thd( - torch.ones(1, 2, device="cuda"), None, None, None, _make_psp([4]), target_len=3 + torch.ones(1, 2, device="cuda"), None, None, None, psp, target_len=3 ) @pytest.mark.internal @@ -915,8 +917,6 @@ def test_reconstruct_preserves_cu_tensors_and_uses_conservative_padding_flag(sel ) } layer = _build_layer(256, 4, 4, 1024, 128, 8) - # Use the non-default mode so losing it during reconstruction is observable. - layer.config.cp_partition_mode = "contiguous" kw = {'packed_seq_params': psp, 'other': 'kept'} TransformerLayer._decompose_packed_seq_params_to_kwargs(kw) assert 'packed_seq_params' not in kw and 'cu_seqlens_q' in kw @@ -924,7 +924,7 @@ def test_reconstruct_preserves_cu_tensors_and_uses_conservative_padding_flag(sel r = kw['packed_seq_params'] assert r.qkv_format == 'thd' and r.max_seqlen_q == 128 assert r.pad_between_seqs is True - assert r.cp_partition_mode == "contiguous" + assert r.cp_partition_mode == "zigzag" for k, v in orig.items(): assert torch.equal(getattr(r, k), v)