diff --git a/megatron/core/distributed/finalize_model_grads.py b/megatron/core/distributed/finalize_model_grads.py index 8cbda1ca0fd..d5b1e77066f 100644 --- a/megatron/core/distributed/finalize_model_grads.py +++ b/megatron/core/distributed/finalize_model_grads.py @@ -467,7 +467,9 @@ def finalize_model_grads( # fp32 gate wgrad 做一次 DP all-reduce, 与参考实现一致。 from ..transformer.module import _use_accuracy_compatible - loss_normalized_in_graph = _use_accuracy_compatible() and (num_tokens is not None) + loss_normalized_in_graph = ( + _use_accuracy_compatible() and not config.dsa_accuracy_compatible and num_tokens is not None + ) if loss_normalized_in_graph: num_tokens = None @@ -573,6 +575,7 @@ def finalize_model_grads( for model_chunk in model: model_chunk.scale_gradients(1.0 / dp_size) + if loss_normalized_in_graph or config.dsa_accuracy_compatible: for model_chunk in model: for param in model_chunk.parameters(): gate_wgrad = getattr(param, "_run_torch_gate_fp32_wgrad", None) diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index b7de1013695..103c2a9a7d6 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -916,8 +916,12 @@ def __init__( for param in self.parameters(): if is_expert: - # Reduce the gradient on the expert_data_parallel group for expert linear layers - setattr(param, "allreduce", not self.expert_parallel) + # Reduce the gradient on the expert_data_parallel group for expert linear layers. + # See _expert_grads_need_own_dp_domain in tensor_parallel/layers.py: ETP < TP + # also puts expert grads in their own (larger) data-parallel domain. + from ..tensor_parallel.layers import _expert_grads_need_own_dp_domain + + setattr(param, "allreduce", not _expert_grads_need_own_dp_domain(self.config)) else: # Reduce the gradient on DP group setattr(param, "allreduce", True) @@ -2035,7 +2039,15 @@ def __init__( ) for param in self.parameters(): - setattr(param, "allreduce", not (is_expert and self.expert_parallel)) + # See _expert_grads_need_own_dp_domain in tensor_parallel/layers.py: + # ETP < TP also puts expert grads in their own data-parallel domain. + from ..tensor_parallel.layers import _expert_grads_need_own_dp_domain + + setattr( + param, + "allreduce", + not (is_expert and _expert_grads_need_own_dp_domain(self.config)), + ) # Explicitly stamp partition_dim and partition_stride on expert weight # tensors when explicit_expert_comm cleared parallel_mode. TE ≤2.12 diff --git a/megatron/core/model_parallel_config.py b/megatron/core/model_parallel_config.py index 88bb070e105..1e1c367d89c 100644 --- a/megatron/core/model_parallel_config.py +++ b/megatron/core/model_parallel_config.py @@ -150,6 +150,9 @@ class ModelParallelConfig: be synchronized. """ + use_accuracy_compatible: bool = False + """Use explicit accuracy-compatible arithmetic in model layers.""" + deterministic_mode: bool = False """If true, code that has deterministic execution will be chosen. This usually means slower execution, but is good for debugging and testing. Defaults to False.""" diff --git a/megatron/core/models/backends.py b/megatron/core/models/backends.py index a270161ddd6..e71db578f14 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -99,6 +99,17 @@ def activation_func(self) -> TEActivationFunctionBuilder | None: class LocalSpecProvider(BackendSpecProvider): """A protocol for providing Local submodules used in Spec building.""" + def linear(self) -> type: + """TP-replicated local Linear (modelopt Linear, not TELinear). + + DSA indexer / MLA down-projections call backend.linear(). TESpecProvider + still returns TELinear; this method is the TE-off counterpart so a + LocalSpecProvider DSA spec does not re-enter Transformer Engine. + """ + from megatron.core.post_training.modelopt.layers import Linear + + return Linear + def column_parallel_linear(self) -> type: """Which column parallel linear module the backend uses""" return ColumnParallelLinear diff --git a/megatron/core/models/common/language_module/language_module.py b/megatron/core/models/common/language_module/language_module.py index fb7b8ef405e..615e2c45196 100644 --- a/megatron/core/models/common/language_module/language_module.py +++ b/megatron/core/models/common/language_module/language_module.py @@ -206,7 +206,7 @@ def compute_language_model_loss(self, labels: Tensor, logits: Tensor) -> Tensor: elif self.config.cross_entropy_fusion_impl == 'native': loss = fused_vocab_parallel_cross_entropy(logits, labels, self.pg_collection.tp) else: - if _use_accuracy_compatible(): + if _use_accuracy_compatible() and not self.config.dsa_accuracy_compatible: s, b = labels.shape loss = torch.nn.functional.cross_entropy( logits.float().reshape(s * b, -1), # [s*b, vocab] 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 a76fe6e3a23..4b01ba561b6 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -20,6 +20,7 @@ ) from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.spec_utils import ModuleSpec +from megatron.core.transformer.torch_norm import WrappedTorchNorm from megatron.core.transformer.transformer_block import ( TransformerBlockSubmodules, get_num_layers_to_build, @@ -57,6 +58,13 @@ ########## +def _get_standalone_norm(config: TransformerConfig, backend: BackendSpecProvider, *, for_qk=False): + rms_norm = config.normalization == "RMSNorm" + if rms_norm and config.norm_accuracy_compatible: + return WrappedTorchNorm + return backend.layer_norm(rms_norm=rms_norm, for_qk=for_qk) + + def get_gated_delta_net_module_spec( config: TransformerConfig, backend: BackendSpecProvider = None ) -> ModuleSpec: @@ -65,12 +73,11 @@ def get_gated_delta_net_module_spec( if backend is None: backend = _get_backend_spec_provider(config=config) - rms_norm = config.normalization == "RMSNorm" attention = ModuleSpec( module=GatedDeltaNet, submodules=GatedDeltaNetSubmodules( in_proj=backend.column_parallel_layer_norm_linear(), - out_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False), + out_norm=_get_standalone_norm(config, backend), out_proj=backend.row_parallel_linear(), ), metainfo={"fuse_input_layernorm": True}, @@ -102,12 +109,10 @@ def get_dsa_module_spec_for_backend( ), ) - # Adjust for RMS norm. - rms_norm = config.normalization == "RMSNorm" # DSA indexer requires normalized q as input, so here we cannot fuse qk layernorm # with linear projection and have to use unfused qk layernorm. qk_norm = ( - backend.layer_norm(rms_norm=rms_norm, for_qk=True) if config.qk_layernorm else IdentityOp + _get_standalone_norm(config, backend, for_qk=True) if config.qk_layernorm else IdentityOp ) attention = ModuleSpec( @@ -228,7 +233,6 @@ def get_transformer_layer_with_experimental_attention_variant_spec( dense_mlp_layer_spec, fuse_layernorm_pre_dense = None, False # Get GPT decoder block layer specs - rms_norm = config.normalization == "RMSNorm" layer_specs = [] for layer_number in range(config.num_layers): attention = ( @@ -245,12 +249,10 @@ def get_transformer_layer_with_experimental_attention_variant_spec( input_layernorm = ( IdentityOp if attention.metainfo["fuse_input_layernorm"] - else backend.layer_norm(rms_norm=rms_norm, for_qk=False) + else _get_standalone_norm(config, backend) ) pre_mlp_layernorm = ( - IdentityOp - if fuse_pre_mlp_layernorm - else backend.layer_norm(rms_norm=rms_norm, for_qk=False) + IdentityOp if fuse_pre_mlp_layernorm else _get_standalone_norm(config, backend) ) layer_specs.append( @@ -317,9 +319,8 @@ def get_transformer_block_with_experimental_attention_variant_spec( layer_specs = [layer_specs[layer_id] for layer_id in local_layer_ids] # Get GPT decoder block spec - rms_norm = config.normalization == "RMSNorm" gpt_decoder_block_spec = TransformerBlockSubmodules( - layer_specs=layer_specs, layer_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False) + layer_specs=layer_specs, layer_norm=_get_standalone_norm(config, backend) ) return gpt_decoder_block_spec diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 984840b3a87..63a1aa51b02 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -773,7 +773,7 @@ def get_gpt_mtp_block_spec_for_backend( raise ValueError(f"Invalid spec: {spec}") mtp_layer_spec = get_mtp_layer_spec_for_backend( - mtp_model_layer_spec=transformer_layer_spec, backend=backend + mtp_model_layer_spec=transformer_layer_spec, backend=backend, config=config ) mtp_num_layers = config.mtp_num_layers if config.mtp_num_layers else 0 if config.mtp_use_repeated_layer: diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 32a61cf7efc..c7406191dc2 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -556,6 +556,9 @@ def _get_megatron_optimizer_based_on_param_groups( # on source of optimizer (Torch or TE/Apex) if USING_PYTORCH_OPTIMIZER: adam_cls = torch.optim.AdamW if config.decoupled_weight_decay else torch.optim.Adam + elif config.native_unfused_adamw: + adam_cls = torch.optim.AdamW if config.decoupled_weight_decay else torch.optim.Adam + kwargs.update({"foreach": False, "fused": False}) else: kwargs["adam_w_mode"] = config.decoupled_weight_decay adam_cls = Adam diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 24f9a032c47..a11f5c8cc3b 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -247,6 +247,9 @@ class OptimizerConfig: adam_eps: float = 1e-08 """Term added to the denominator to improve numerical stability in Adam optimizer.""" + native_unfused_adamw: bool = False + """Use torch.optim.AdamW with foreach=False and fused=False instead of TE/Apex Adam.""" + decoupled_weight_decay: bool = True """If true, decouples weight decay from the gradient update, equivalent to AdamW. If false, original Adam update rule will be used. Defaults to True. diff --git a/megatron/core/pipeline_parallel/schedules.py b/megatron/core/pipeline_parallel/schedules.py index e67c498e2cc..f576bc05a58 100644 --- a/megatron/core/pipeline_parallel/schedules.py +++ b/megatron/core/pipeline_parallel/schedules.py @@ -163,7 +163,7 @@ def forward_step(data_iterator, model): return forward_backward_func -def deallocate_output_tensor(out, deallocate_pipeline_outputs=False): +def deallocate_output_tensor(out, deallocate_pipeline_outputs=False, config=None): '''Pseudo-deallocate (i.e., set to scalar) the output tensor's '.data' field. This method should be called right after the output tensor has been @@ -177,17 +177,23 @@ def deallocate_output_tensor(out, deallocate_pipeline_outputs=False): ''' if (out is None) or (not deallocate_pipeline_outputs): return + if ( + config is not None + and config.dsa_accuracy_compatible + and config.tensor_model_parallel_size <= 1 + ): + return # Handle dict format (multi-module pipelines) if isinstance(out, dict): for value in out.values(): - deallocate_output_tensor(value, deallocate_pipeline_outputs) + deallocate_output_tensor(value, deallocate_pipeline_outputs, config) return # Handle list format if isinstance(out, list): for item in out: - deallocate_output_tensor(item, deallocate_pipeline_outputs) + deallocate_output_tensor(item, deallocate_pipeline_outputs, config) return # Base case: deallocate tensor @@ -568,7 +574,10 @@ def backward_step(input_tensor, output_tensor, output_tensor_grad, config): # This results in a tensor that does not require gradients. # In such cases, we intentionally skip the backward pass while preserving zero gradients. if output_tensor[0].requires_grad: - if config.deallocate_pipeline_outputs: + _tp_size = int(getattr(config, "tensor_model_parallel_size", 1) or 1) + if config.deallocate_pipeline_outputs and ( + not config.dsa_accuracy_compatible or _tp_size > 1 + ): custom_backward(output_tensor[0], output_tensor_grad[0]) else: torch.autograd.backward(output_tensor[0], grad_tensors=output_tensor_grad[0]) @@ -641,7 +650,10 @@ def _unwrap_single_tensor_list(tensor): # In multi-modal models like VLM, some batches may not have images. # In such cases, skip backward while preserving zero gradients. if output_tensor_module is not None and output_tensor_module.requires_grad: - if config.deallocate_pipeline_outputs: + _tp_size = int(getattr(config, "tensor_model_parallel_size", 1) or 1) + if config.deallocate_pipeline_outputs and ( + not config.dsa_accuracy_compatible or _tp_size > 1 + ): custom_backward(output_tensor_module, output_tensor_grad_module) else: torch.autograd.backward( @@ -1637,7 +1649,7 @@ def forward_backward_helper_wrapper( ) if recv_prev: input_tensors[next_forward_model_chunk_id].append(input_tensor) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) else: if not is_pp_first_stage(p2p_communicator.pp_group): # Send only since recv prefetched. @@ -1667,7 +1679,7 @@ def forward_backward_helper_wrapper( send_next_wait_handle.wait() send_next_wait_handle = None - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) if recv_prev: input_tensors[next_forward_model_chunk_id].append( fwd_recv_buffer[k % fwd_recv_buffer_size] @@ -1745,7 +1757,7 @@ def pp_pre_forward(vp_stage=None): recv_prev_wait_handle = recv_prev_wait_handles.pop(0) recv_prev_wait_handle.wait() - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Async forward send / receive def pp_post_forward(output_tensor, vp_stage=None): @@ -1920,7 +1932,7 @@ def pp_post_backward(input_tensor_grad, vp_stage=None): tensor_shape=tensor_shape, ) ) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Put input_tensor and output_tensor_grad in data structures in the # right location. if recv_prev: @@ -1928,7 +1940,7 @@ def pp_post_backward(input_tensor_grad, vp_stage=None): if recv_next: output_tensor_grads[next_backward_model_chunk_id].append(output_tensor_grad) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) nvtx_range_pop(suffix="steady") # Run cooldown backward passes (flush out pipeline) for the last model chunk. @@ -2353,7 +2365,7 @@ def enable_grad_sync(): if not forward_only: input_tensors.append(input_tensor) output_tensors.append(output_tensor) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Before running 1F1B, need to receive first forward tensor. # If all microbatches are run in warmup / cooldown phase, then no need to @@ -2408,7 +2420,7 @@ def enable_grad_sync(): # Add input_tensor and output_tensor to end of list. input_tensors.append(input_tensor) output_tensors.append(output_tensor) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Pop input_tensor and output_tensor from the start of the list for # the backward pass. diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 2f927d5218d..b613c581ed2 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -59,6 +59,51 @@ from megatron.core.transformer.module import _use_accuracy_compatible + +class _EmbedFp32MainGrad(torch.autograd.Function): + """UAC embedding lookup whose wgrad lands in fp32 main_grad. + + Forward is ``weight[ids]`` (bf16 activation unchanged). Backward uses a + unique-row clone plus ``autograd.grad``, then ``index_add_`` into an fp32 + ``main_grad``. Returns None for weight.grad so MixPrecision cannot add_(bf16). + """ + + @staticmethod + def forward(ctx, weight, ids): + """Look up embeddings while retaining the accumulator owner.""" + ctx.save_for_backward(ids) + ctx.weight_ref = weight + return weight[ids] + + @staticmethod + def backward(ctx, grad_output): + """Accumulate repeated-index gradients into the FP32 master buffer.""" + (ids,) = ctx.saved_tensors + weight = ctx.weight_ref + prev = torch.is_grad_enabled() + torch.set_grad_enabled(True) + try: + ids_flat = ids.reshape(-1) + unique_ids, inv = torch.unique(ids_flat, return_inverse=True) + uniq_w = weight.detach()[unique_ids].clone().requires_grad_(True) + looked = uniq_w[inv.reshape(ids.shape)] + (gw,) = torch.autograd.grad(looked, uniq_w, grad_outputs=grad_output, allow_unused=True) + finally: + torch.set_grad_enabled(prev) + if gw is None: + return None, None + fp = gw.float() + if hasattr(weight, "main_grad") and weight.main_grad is not None: + weight.main_grad.index_add_(0, unique_ids, fp) + else: + acc = torch.zeros(weight.shape, dtype=torch.float32, device=weight.device) + acc.index_add_(0, unique_ids, fp) + weight.main_grad = acc + if hasattr(weight, "grad_added_to_main_grad"): + weight.grad_added_to_main_grad = True + return None, None + + _MODEL_PARALLEL_ATTRIBUTE_DEFAULTS = { "expert_tp": False, "is_qkv": False, @@ -237,7 +282,11 @@ def __init__( ) ) self.num_embeddings_per_partition = self.vocab_end_index - self.vocab_start_index - self.deterministic_mode = config.deterministic_mode or _use_accuracy_compatible() + self.deterministic_mode = ( + config.deterministic_mode + or _use_accuracy_compatible() + or config.use_accuracy_compatible + ) self.config = config self.use_inference_optimized_reduce_scatter = ( @@ -299,7 +348,11 @@ def forward(self, input_): masked_input = input_ # Get the embeddings. if self.deterministic_mode: - output_parallel = self.weight[masked_input] + _tp_size = 1 if self.tp_group is None else self.tp_group.size() + if getattr(self.config, "dsa_accuracy_compatible", False) and _tp_size <= 1: + output_parallel = _EmbedFp32MainGrad.apply(self.weight, masked_input) + else: + output_parallel = self.weight[masked_input] else: # F.embedding currently has a non-deterministic backward function output_parallel = F.embedding(masked_input, self.weight) @@ -675,6 +728,7 @@ def linear_with_grad_accumulation_and_async_allreduce( grad_output_buffer: Optional[List[torch.Tensor]] = None, wgrad_deferral_limit: Optional[int] = 0, tp_group: Optional[torch.distributed.ProcessGroup] = None, + dsa_accuracy_compatible: bool = False, ) -> torch.Tensor: """Linear layer execution with asynchronous communication and gradient accumulation fusion in backprop. @@ -740,6 +794,12 @@ def linear_with_grad_accumulation_and_async_allreduce( """ tp_group = get_tensor_model_parallel_group_if_none(tp_group) + _tp_size = 1 if tp_group is None else tp_group.size() + if dsa_accuracy_compatible and _tp_size <= 1 and not sequence_parallel and not allreduce_dgrad: + output = torch.matmul(input, weight.t()) + if bias is not None: + output = output + bias + return output args = [ input, @@ -777,6 +837,23 @@ def linear_with_grad_accumulation_and_async_allreduce( linear_with_grad_accumulation_and_async_allreduce.warned = False +def _expert_grads_need_own_dp_domain(config) -> bool: + """Use expert DP for split expert tensor groups in DSA alignment. + + The default retains EP-only grouping. With DSA TP2/ETP1, expert replicas + consume different sequence shards: reduce their gradients over expt_dp + instead of the dense dp_cp group. + """ + if config.expert_model_parallel_size > 1: + return True + if not getattr(config, 'dsa_accuracy_compatible', False): + return False + etp = getattr(config, 'expert_tensor_parallel_size', None) + if etp is None: + return False + return etp != config.tensor_model_parallel_size + + class ColumnParallelLinear(torch.nn.Module): """Linear layer with column parallelism. @@ -921,7 +998,11 @@ def __init__( tensor=self.weight, is_parallel=True, dim=0, stride=stride ) - setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.weight, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) else: self.weight = None @@ -943,7 +1024,11 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.bias, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) else: self.register_parameter("bias", None) @@ -991,7 +1076,13 @@ def _forward_impl(self, input, weight, *args, **kwargs): if not weight.requires_grad: return linear_with_frozen_weight(input, weight, *args, **kwargs) else: - return linear_with_grad_accumulation_and_async_allreduce(input, weight, *args, **kwargs) + return linear_with_grad_accumulation_and_async_allreduce( + input, + weight, + *args, + dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), + **kwargs, + ) def forward( self, @@ -1039,6 +1130,10 @@ def forward( or self.disable_grad_reduce ): input_parallel = input_ + elif getattr(self.config, "dsa_accuracy_compatible", False) and ( + self.tp_group is None or self.tp_group.size() <= 1 + ): + input_parallel = input_ else: input_parallel = copy_to_tensor_model_parallel_region(input_, group=self.tp_group) @@ -1084,7 +1179,10 @@ def forward( if runtime_gather_output is not None: gather_output = runtime_gather_output - if gather_output: + if gather_output and ( + not getattr(self.config, "dsa_accuracy_compatible", False) + or (self.tp_group is not None and self.tp_group.size() > 1) + ): # All-gather across the partitions. if self.use_inference_optimized_all_gather and not self.training: # Deferred to avoid circular import: inference_layers → TE → layers. @@ -1271,7 +1369,11 @@ def __init__( set_tensor_model_parallel_attributes( tensor=self.weight, is_parallel=True, dim=1, stride=stride ) - setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.weight, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) if bias: if config.use_cpu_initialization: @@ -1289,7 +1391,11 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.bias, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) setattr(self.bias, "sequence_parallel", self.sequence_parallel) else: self.register_parameter("bias", None) @@ -1305,7 +1411,13 @@ def _forward_impl(self, input, weight, *args, **kwargs): if not weight.requires_grad: return linear_with_frozen_weight(input, weight, *args, **kwargs) else: - return linear_with_grad_accumulation_and_async_allreduce(input, weight, *args, **kwargs) + return linear_with_grad_accumulation_and_async_allreduce( + input, + weight, + *args, + dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), + **kwargs, + ) def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Forward of RowParallelLinear @@ -1355,6 +1467,10 @@ def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: output_ = reduce_scatter_to_sequence_parallel_region( output_parallel, group=self.tp_group ) + elif getattr(self.config, "dsa_accuracy_compatible", False) and ( + self.tp_group is None or self.tp_group.size() <= 1 + ): + output_ = output_parallel else: output_ = reduce_from_tensor_model_parallel_region(output_parallel, group=self.tp_group) if not self.skip_bias_add: diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index dde238635c2..8e0c135d725 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -64,6 +64,7 @@ def _unfused_absorbed_dsa_fn( varlen_starts: Optional[torch.Tensor] = None, varlen_ends: Optional[torch.Tensor] = None, key_positions: Optional[torch.Tensor] = None, + accuracy_compatible: bool = False, ) -> torch.Tensor: """Unfused absorbed-MLA attention: output stays [sq, b, np, v_channels].""" sq, b, np, hn = query.size() @@ -99,10 +100,15 @@ def _unfused_absorbed_dsa_fn( ) attention_scores = attention_scores + index_mask.unsqueeze(1) - valid_index_mask = torch.isfinite(index_mask) - attention_scores = dsa_masking.masked_softmax( - attention_scores.float(), valid_index_mask.unsqueeze(1).expand(b, np, sq, skv), dim=-1 - ) + valid_index_mask = torch.isfinite(index_mask).unsqueeze(1).expand(b, np, sq, skv) + if accuracy_compatible: + attention_scores = _AccuracyCompatibleSoftmax.apply( + attention_scores.float(), valid_index_mask + ) + else: + attention_scores = dsa_masking.masked_softmax( + attention_scores.float(), valid_index_mask, dim=-1 + ) # Latent value is the first v_channels slice of absorbed key cache. value = key[..., :v_channels].permute(1, 2, 0, 3) # [b,1,skv,v] @@ -110,6 +116,25 @@ def _unfused_absorbed_dsa_fn( return output.permute(2, 0, 1, 3).contiguous() +class _AccuracyCompatibleSoftmax(torch.autograd.Function): + """Masked softmax with an explicit backward formula for DSA alignment.""" + + @staticmethod + def forward(ctx, logits: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor: + probabilities = torch.softmax(logits.masked_fill(~valid_mask, float("-inf")), dim=-1) + probabilities = probabilities.masked_fill(~valid_mask, 0.0) + ctx.save_for_backward(probabilities, valid_mask) + return probabilities + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + probabilities, valid_mask = ctx.saved_tensors + grad_logits = probabilities * ( + grad_output - (grad_output * probabilities).sum(dim=-1, keepdim=True) + ) + return grad_logits.masked_fill(~valid_mask, 0.0), None + + def _run_sparse_attention( *, absorbed_mla: bool, @@ -127,6 +152,7 @@ def _run_sparse_attention( topk_length: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Run sparse attention for absorbed and non-absorbed MLA paths.""" + accuracy_compatible = bool(getattr(config, "dsa_accuracy_compatible", False)) if absorbed_mla: latent_v_channels = int(getattr(config, "kv_lora_rank", 0) or 0) if latent_v_channels <= 0: @@ -143,7 +169,7 @@ def _run_sparse_attention( "Received absorbed layout with explicit value tensor." ) output = None - if dsa_kernels.use_fused_dsa_kernels(config): + if not accuracy_compatible and dsa_kernels.use_fused_dsa_kernels(config): output = dsa_kernels.run_fused_absorbed_sparse_attention( config, query, @@ -166,6 +192,7 @@ def _run_sparse_attention( varlen_starts=varlen_starts, varlen_ends=varlen_ends, key_positions=key_positions, + accuracy_compatible=accuracy_compatible, ) assert output is not None output = torch.einsum("sbhc,hdc->sbhd", output, up_v_weight).contiguous() @@ -182,6 +209,7 @@ def _run_sparse_attention( varlen_starts=varlen_starts, varlen_ends=varlen_ends, key_positions=key_positions, + accuracy_compatible=accuracy_compatible, ) @@ -1411,6 +1439,7 @@ def unfused_dsa_fn( varlen_starts: Optional[torch.Tensor] = None, varlen_ends: Optional[torch.Tensor] = None, key_positions: Optional[torch.Tensor] = None, + accuracy_compatible: bool = False, ): """ Unfused sparse attention implementation. @@ -1457,6 +1486,27 @@ def unfused_dsa_fn( device=query.device, ) + if accuracy_compatible: + index_mask = torch.full((b, sq, skv), float("-inf"), device=query.device) + dsa_masking.scatter_topk_into_index_mask(index_mask, topk_indices) + index_mask = dsa_masking.apply_sparse_validity_to_index_mask( + index_mask, + row_mask=row_mask, + varlen_starts=varlen_starts, + varlen_ends=varlen_ends, + key_positions=key_positions, + ) + valid_index_mask = torch.isfinite(index_mask).unsqueeze(1).expand(b, np, sq, skv) + attention_scores = ( + torch.matmul(query_b.float(), key_b.float().transpose(-1, -2)) * softmax_scale + ) + attention_probs = _AccuracyCompatibleSoftmax.apply( + attention_scores + index_mask.unsqueeze(1), valid_index_mask + ) + output = torch.matmul(attention_probs.to(value_b.dtype), value_b) + output = output.permute(2, 0, 1, 3).contiguous().view(sq, b, np * hnv) + return output.squeeze(1) if query_was_thd else output + seq_chunk_size = 512 head_chunk_size = 16 topk_chunk_size = 1024 @@ -1777,8 +1827,11 @@ def forward( skv = key.size(0) # Detach x and qr to prevent gradients of indexer from flowing back to the main model. - x = x.detach() - qr = qr.detach() + _tp_group = getattr(self.pg_collection, "tp", None) + _tp_size = 1 if _tp_group is None else _tp_group.size() + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): + x = x.detach() + qr = qr.detach() indexer_loss_coeff = self.config.dsa_indexer_loss_coeff or 0.0 computes_topk = not self.skip_topk diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 229d02a00b6..67fc3106187 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -1347,35 +1347,24 @@ def forward( output_local_list = [] - for _ei, (expert, tokens, probs) in enumerate( - zip(self.local_experts, tokens_list, probs_list) - ): - # Keep the expert GEMM shape identical to Paddle in bit-exact mode. - # Padding only Paddle aligns tiny-M forward, but changes its backward - # dgrad GEMM from M<17 to M=32. Padding both sides aligns all GEMMs. + for expert, tokens, probs in zip(self.local_experts, tokens_list, probs_list): + # The unfused Paddle expert pads tiny GEMMs to 32 rows. The + # grouped-storage fallback uses real token counts instead. num_real_tokens = tokens.shape[0] pad_small_expert = _use_accuracy_compatible() and 0 < num_real_tokens < 17 + if self.config.dsa_accuracy_compatible: + pad_small_expert = ( + self.config.use_accuracy_compatible + and not self.config.moe_grouped_gemm + and not (self.config.fp8 or self.config.fp4) + and 0 < num_real_tokens < 17 + ) if pad_small_expert: num_pad_tokens = 32 - num_real_tokens tokens = torch.cat( - ( - tokens, - torch.zeros( - num_pad_tokens, - tokens.shape[1], - dtype=tokens.dtype, - device=tokens.device, - ), - ), - dim=0, - ) - probs = torch.cat( - ( - probs, - torch.zeros(num_pad_tokens, dtype=probs.dtype, device=probs.device), - ), - dim=0, + (tokens, tokens.new_zeros(num_pad_tokens, tokens.shape[1])), dim=0 ) + probs = torch.cat((probs, probs.new_zeros(num_pad_tokens)), dim=0) if self.config.fp8 or self.config.fp4: hidden, probs = self._pad_tensor_for_quantization(tokens, probs) output, output_bias = expert(hidden, probs) diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index 5eb8de35f09..798eb4cabd6 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -576,7 +576,11 @@ def postprocess(self, output: torch.Tensor, shared_expert_output: Optional[torch output, _ = self.fc2_latent_proj(output) if shared_expert_output is not None: - output = output + shared_expert_output + if self.config.dsa_accuracy_compatible: + orig_dtype = output.dtype + output = (output.float() + shared_expert_output.float()).to(orig_dtype) + else: + output = output + shared_expert_output elif ( isinstance(self.token_dispatcher, NVLSAllGatherVDispatcher) and self._latent_shared_expert_output is not None @@ -660,7 +664,11 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): hidden_states_router = hidden_states hidden_states_dispatch = hidden_states - shared_expert_output = self.shared_experts_compute(hidden_states_shared) + if self.config.dsa_accuracy_compatible and not self.shared_expert_overlap: + self._accuracy_shared_input = hidden_states_shared + shared_expert_output = None + else: + shared_expert_output = self.shared_experts_compute(hidden_states_shared) probs, routing_map = self.route(hidden_states_router, padding_mask) hidden_states, probs = self.preprocess( hidden_states_dispatch, probs, routing_map @@ -695,6 +703,12 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): if intermediate_tensors is not None: output, shared_expert_output = intermediate_tensors + if self.config.dsa_accuracy_compatible: + shared_input = getattr(self, "_accuracy_shared_input", None) + if shared_input is not None: + shared_expert_output = self.shared_experts_compute(shared_input) + self._accuracy_shared_input = None + output = self.postprocess(output, shared_expert_output) if intermediate_tensors is not None: diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index afea9dcbf34..6da367bb659 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -404,6 +404,7 @@ def permute( drop_and_pad: bool = False, tokens_per_expert: Optional[torch.Tensor] = None, align_size: int = 0, + dsa_accuracy_compatible: bool = False, ) -> Tuple[ torch.Tensor, Optional[torch.Tensor], @@ -524,7 +525,11 @@ def permute( permuted_probs = probs.T.contiguous().reshape(-1)[flat_sorted] # === BIT-EXACT permute backward (gated by MOE_DETERMINISTIC_UNPERMUTE) === - if _use_accuracy_compatible() and not (drop_and_pad and num_out_tokens is not None): + if ( + _use_accuracy_compatible() + and not dsa_accuracy_compatible + and not (drop_and_pad and num_out_tokens is not None) + ): rm_T_int = routing_map.long() # [num_experts, num_tokens] tokens_per_expert_local = rm_T_int.sum(dim=-1) # [num_experts] expert_offsets = torch.zeros(num_experts + 1, dtype=torch.long, device=tokens.device) diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index c9aedd860c6..a1ca6cdbe42 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -104,6 +104,12 @@ def gating(self, input: torch.Tensor): router_dtype = torch.float32 elif self.config.moe_router_dtype == 'fp64': router_dtype = torch.float64 + if self.config.router_accuracy_compatible: + inp_shape = input.shape + logits = torch.mm(input.reshape(-1, inp_shape[-1]).float(), self.weight.float().t()) + if self.bias is not None: + logits = logits + self.bias.float() + return logits.view(*inp_shape[:-1], -1) logits = router_gating_linear(input, self.weight, self.bias, router_dtype) return logits diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 490e4d8cc6d..1b26b87e4e4 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -304,6 +304,7 @@ def dispatch_postprocess(self, hidden_states, probs): self.local_map, num_out_tokens=tokens_per_expert.sum().item(), fused=self.config.moe_permute_fusion, + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) ) @@ -651,6 +652,7 @@ def dispatch_preprocess( num_out_tokens=self.num_out_tokens, fused=self.config.moe_permute_fusion, drop_and_pad=self.drop_and_pad, + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) return permutated_local_input_tokens, permuted_probs @@ -1377,6 +1379,7 @@ def get_permuted_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> fused=self.permute_fusion, tokens_per_expert=self.tokens_per_expert, align_size=get_align_size_for_quantization(self.config), + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) if self.router_dtype == "fp64": permuted_probs = permuted_probs.to(torch.float64) @@ -1526,6 +1529,9 @@ def token_dispatch( """ if self.shared_experts is not None: self.shared_experts.wait_current_stream() + if self.config.dsa_accuracy_compatible: + async_finish = False + allocate_on_comm_stream = False dispatched_hidden_states = self._comm_manager.dispatch( hidden_states, async_finish, allocate_on_comm_stream ) @@ -1585,6 +1591,9 @@ def token_combine( # when CUDA_DEVICE_MAX_CONNECTIONS>1. if self.shared_experts is not None: self.shared_experts.wait_current_stream() + if self.config.dsa_accuracy_compatible: + async_finish = False + allocate_on_comm_stream = False return self._comm_manager.combine(hidden_states, async_finish, allocate_on_comm_stream) def combine_postprocess(self, hidden_states: torch.Tensor): diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index b20514ce6a4..5a91a7da037 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -30,7 +30,7 @@ from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module -from megatron.core.transformer.torch_norm import LayerNormBuilder +from megatron.core.transformer.torch_norm import LayerNormBuilder, WrappedTorchNorm from megatron.core.transformer.transformer_block import TransformerBlockSubmodules from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.typed_torch import apply_module @@ -577,7 +577,9 @@ class MultiTokenPredictionLayerSubmodules: def get_mtp_layer_spec( - mtp_model_layer_spec: ModuleSpec, use_transformer_engine: bool + mtp_model_layer_spec: ModuleSpec, + use_transformer_engine: bool, + config: Optional[TransformerConfig] = None, ) -> ModuleSpec: """Get the MTP layer spec. @@ -587,11 +589,14 @@ def get_mtp_layer_spec( return get_mtp_layer_spec_for_backend( mtp_model_layer_spec, backend=TESpecProvider() if use_transformer_engine else LocalSpecProvider(), + config=config, ) def get_mtp_layer_spec_for_backend( - mtp_model_layer_spec: ModuleSpec, backend: BackendSpecProvider + mtp_model_layer_spec: ModuleSpec, + backend: BackendSpecProvider, + config: Optional[TransformerConfig] = None, ) -> ModuleSpec: """Get the MTP layer spec. @@ -599,7 +604,11 @@ def get_mtp_layer_spec_for_backend( ModuleSpec: Module specification with modules from the backend. """ column_parallel_linear_impl: type = backend.column_parallel_linear() - layer_norm_impl = backend.layer_norm() + layer_norm_impl = ( + WrappedTorchNorm + if config is not None and config.norm_accuracy_compatible + else backend.layer_norm() + ) mtp_layer_spec = ModuleSpec( module=MultiTokenPredictionLayer, submodules=MultiTokenPredictionLayerSubmodules( @@ -1106,7 +1115,11 @@ def _get_embeddings( if self.config.mtp_detach_heads: decoder_input = decoder_input.detach() - hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + _tp_size = 1 if self.tp_group is None else self.tp_group.size() + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # make_viewless_tensor no-ops when hidden_states is not a view (_base is None), # which happens after detach() with mtp_detach_heads. Activation # checkpointing (CheckpointFunction.apply) requires at least one input tensor @@ -1121,10 +1134,17 @@ def _concat_embeddings(self, hidden_states: torch.Tensor, decoder_input: torch.T """ Concatenate the tokens before sending to transformer layer. """ + _tp_size = 1 if self.tp_group is None else self.tp_group.size() decoder_input = apply_module(self.enorm)(decoder_input) - decoder_input = make_viewless_tensor(inp=decoder_input, requires_grad=True, keep_graph=True) + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): + decoder_input = make_viewless_tensor( + inp=decoder_input, requires_grad=True, keep_graph=True + ) hidden_states = apply_module(self.hnorm)(hidden_states) - hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # At the (k - 1)-th MTP module, concatenates the i-th token's hidden_states # and the (i + K)-th token's embedding, and combine them with linear projection. hidden_states = torch.cat((decoder_input, hidden_states), -1) @@ -1135,7 +1155,7 @@ def _concat_embeddings(self, hidden_states: torch.Tensor, decoder_input: torch.T hidden_states = inference_all_gather_from_tensor_model_parallel_region( hidden_states, self.tp_group, self.config ) - else: + elif not (self.config.dsa_accuracy_compatible and _tp_size <= 1): hidden_states = gather_from_tensor_model_parallel_region( hidden_states, group=self.tp_group ) diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py index 5948ae600f9..a7f9db5e60a 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -47,7 +47,9 @@ def __new__( assert not config.persist_layer_norm, f"persist_layer_norm not supported by torch LayerNorm" - assert not config.sequence_parallel, f"sequence parallel not supported by torch LayerNorm" + assert ( + config.norm_accuracy_compatible or not config.sequence_parallel + ), "sequence parallel not supported by torch LayerNorm" assert ( not config.memory_efficient_layer_norm @@ -66,7 +68,14 @@ def __new__( else: raise Exception("Only LayerNorm, RMSNorm and L2Norm are currently supported") - return norm_cls(normalized_shape=hidden_size, eps=eps) + factory_kwargs = {} + if config.normalization == "RMSNorm" and config.norm_accuracy_compatible: + factory_kwargs["dtype"] = config.params_dtype + norm = norm_cls(normalized_shape=hidden_size, eps=eps, **factory_kwargs) + if config.sequence_parallel: + for parameter in norm.parameters(): + parameter.sequence_parallel = True + return norm class L2Norm(torch.nn.Module, LayerNormInterface): diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index 0415035ffbe..b3b977f8053 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -590,7 +590,11 @@ def forward( # likely redundant, since p2p_communication.py (likely originator) # already creates viewless tensors. That said, make_viewless_tensor() # is called here to be future-proof and corner-case-proof. - hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + _tp_size = int(getattr(self.config, "tensor_model_parallel_size", 1) or 1) + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) if self.config.sequence_parallel: rng_context = tensor_parallel.get_cuda_rng_tracker().fork() @@ -694,9 +698,11 @@ def forward( # TENorm produces a "viewed" tensor. This will result in schedule.py's # deallocate_output_tensor() throwing an error, so a viewless tensor is # created to prevent this. - hidden_states = make_viewless_tensor( - inp=hidden_states, requires_grad=True, keep_graph=True - ) + _tp_size = int(getattr(self.config, "tensor_model_parallel_size", 1) or 1) + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # If this TransformerBlock is empty, input and output hidden states will be the same node # on the computational graph and will lead to unexpected errors in pipeline schedules. diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index bbcf413baee..34c02443365 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -186,6 +186,16 @@ class TransformerConfig(ModelParallelConfig): ) """Epsilon value for any LayerNorm/RMSNorm operations.""" + norm_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--norm-accuracy-compatible"]}} + ) + """Use native Torch RMSNorm modules instead of Transformer Engine norm modules for alignment.""" + + router_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--router-accuracy-compatible"]}} + ) + """Use an explicit fp32 router GEMM instead of the fused Transformer Engine path.""" + layernorm_zero_centered_gamma: bool = field( default=False, metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}} ) @@ -268,6 +278,13 @@ class TransformerConfig(ModelParallelConfig): """Whether cross entropy loss is calculated over the actual number of non-padded tokens in the global batch, versus the default behavior of assuming all tokens are non-padded.""" + accuracy_compatible_loss_sum_dtype: Literal["float32", "float64"] = "float64" + """Token-loss accumulation dtype in accuracy-compatible training. + + Preserve FP64 accumulation by default. Model providers can select FP32 + when that is the reference loss-reduction contract. + """ + multi_latent_attention: bool = False """Whether to use multi-latent attention.""" @@ -318,6 +335,14 @@ class TransformerConfig(ModelParallelConfig): ``none`` disables fused DSA kernels. Explicit ``tilelang`` or ``cudnn`` enables only that backend. Unsupported DSA layouts continue to use the PyTorch fallback.""" + dsa_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--dsa-accuracy-compatible"]}} + ) + """Use DSA reference numerics: explicit softmax backward, deferred token-loss + normalization, FP32 MoE accumulation and TP1 autograd paths. Disabled by + default to preserve existing models' accuracy-compatible behavior. + """ + dsa_indexer_rope_interleaved: bool = False """Whether DSA indexer RoPE should use MLA-style interleaving.""" diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index 904912c18d8..5a584e59f42 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -941,9 +941,12 @@ def _forward_post_mlp( # won't result in memory savings (like the data loader, or # p2p_communication), it serves to document the origin of this # 'view' tensor. - output = make_viewless_tensor( - inp=hidden_states, requires_grad=hidden_states.requires_grad, keep_graph=True - ) + if self.config.dsa_accuracy_compatible and self.config.tensor_model_parallel_size <= 1: + output = hidden_states + else: + output = make_viewless_tensor( + inp=hidden_states, requires_grad=hidden_states.requires_grad, keep_graph=True + ) return output diff --git a/tests/unit_tests/distributed/test_finalize_model_grads.py b/tests/unit_tests/distributed/test_finalize_model_grads.py index ee535c29baf..762e1fdd81d 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -1,5 +1,4 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. -import inspect import os import pytest @@ -21,6 +20,54 @@ from tests.unit_tests.test_utilities import Utils +@pytest.mark.parametrize( + "accuracy,dsa,expected", [(False, False, 1.0), (True, False, 4.0), (True, True, 1.0)] +) +def test_token_normalization_preserves_legacy_and_dsa_contracts( + monkeypatch, accuracy, dsa, expected +): + """Legacy in-graph normalization averages DP; DSA divides by global tokens.""" + import importlib + from types import SimpleNamespace + + implementation = importlib.import_module("megatron.core.distributed.finalize_model_grads") + module = importlib.import_module("megatron.core.transformer.module") + monkeypatch.setattr(module, "_use_accuracy_compatible", lambda: accuracy) + config = TransformerConfig( + num_layers=1, hidden_size=8, num_attention_heads=1, dsa_accuracy_compatible=dsa + ) + gradient = torch.tensor(8.0, device="cuda") + events = [] + + def scale(value): + events.append("scale") + gradient.mul_(value) + + model = SimpleNamespace( + config=config, + parameters=lambda: (), + scale_gradients=scale, + finish_grad_sync=lambda **kwargs: events.append("sync"), + ) + monkeypatch.setattr(implementation, "get_model_config", lambda model: model.config) + for name in ( + "_allreduce_conditional_embedding_grads", + "_allreduce_non_tensor_model_parallel_grads", + "_allreduce_word_embedding_grads", + "_allreduce_position_embedding_grads", + "reset_model_temporary_tensors", + ): + monkeypatch.setattr(implementation, name, lambda *args: None) + monkeypatch.setattr(parallel_state, "get_data_parallel_world_size", lambda **kwargs: 2) + monkeypatch.setattr(implementation, "get_pp_last_rank", lambda group: 0) + monkeypatch.setattr(dist, "broadcast", lambda tensor, **kwargs: None) + monkeypatch.setattr(dist, "all_reduce", lambda tensor, **kwargs: tensor.mul_(2)) + groups = SimpleNamespace(tp=None, pp=None, embd=None, pos_embd=None, dp_cp=None) + finalize_model_grads([model], num_tokens=torch.tensor(4.0, device="cuda"), pg_collection=groups) + assert gradient.item() == expected + assert events == ["sync", "scale"] + + class _RouterExpertBiasModel(torch.nn.Module): def __init__(self, config, local_tokens_per_expert): super().__init__() diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py index 0a454b5d7ff..13cc7419c45 100644 --- a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -65,6 +65,7 @@ def _make_config(**overrides): defaults = dict( num_layers=4, normalization="RMSNorm", + norm_accuracy_compatible=False, qk_layernorm=False, multi_latent_attention=False, qk_l2_norm=False, @@ -369,6 +370,23 @@ def test_qk_layernorm_enabled(self, normalization): assert spec.submodules.q_layernorm is spec.submodules.kv_layernorm backend.layer_norm.assert_any_call(rms_norm=expected_rms, for_qk=True) + def test_accuracy_compatible_qk_rmsnorm(self): + """Verify DSA q/kv norms can use the native Torch RMSNorm builder.""" + from megatron.core.transformer.torch_norm import WrappedTorchNorm + + backend = _make_backend() + cfg = _make_config( + multi_latent_attention=True, + qk_l2_norm=False, + qk_layernorm=True, + normalization="RMSNorm", + norm_accuracy_compatible=True, + ) + spec = self._call(cfg=cfg, backend=backend) + + assert spec.submodules.q_layernorm is WrappedTorchNorm + assert spec.submodules.kv_layernorm is WrappedTorchNorm + def test_qk_layernorm_disabled(self): """Verify q/kv layernorm becomes IdentityOp, skipping backend.layer_norm for qk.""" backend = _make_backend() diff --git a/tests/unit_tests/models/test_local_spec_provider_linear.py b/tests/unit_tests/models/test_local_spec_provider_linear.py new file mode 100644 index 00000000000..4564ce0adf8 --- /dev/null +++ b/tests/unit_tests/models/test_local_spec_provider_linear.py @@ -0,0 +1,15 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +"""LocalSpecProvider must expose a non-TE backend.linear() for DSA/MLA.""" + +from megatron.core.extensions.transformer_engine import TELinear +from megatron.core.models.backends import LocalSpecProvider +from megatron.core.post_training.modelopt.layers import Linear +from megatron.core.tensor_parallel.layers import ColumnParallelLinear + + +def test_local_spec_provider_linear_is_replicated_local_linear(): + backend = LocalSpecProvider() + assert backend.linear() is Linear + assert backend.linear() is not TELinear + assert backend.column_parallel_linear() is ColumnParallelLinear + assert backend.linear() is not backend.column_parallel_linear() diff --git a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py new file mode 100644 index 00000000000..979fff5282f --- /dev/null +++ b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py @@ -0,0 +1,384 @@ +"""CUDA unit tests for native C2 TP1 accuracy-compatible migration.""" + +from __future__ import annotations + +import ast +import os +import unittest +from pathlib import Path +from types import SimpleNamespace +from typing import List, Optional +from unittest.mock import patch + +import torch +import torch.nn.functional as F + +ROOT = Path(__file__).resolve().parents[3] +_UAC = {"on": False} +_TP = {"size": 1} +_CUSTOM_BWD = {"calls": []} + + +def _use_accuracy_compatible(): + return _UAC["on"] + + +def _custom_backward(output, grad): + _CUSTOM_BWD["calls"].append((output, grad)) + output.backward(grad) + + +class _SentinelApply(torch.autograd.Function): + last = None + + @staticmethod + def forward(ctx, *args): + _SentinelApply.last = args + inp = args[0] + return inp.new_zeros(inp.shape[:-1] + (args[1].shape[0],)) + + @staticmethod + def backward(ctx, grad_output): + return (grad_output,) + (None,) * 8 + + +class _FakeGroup: + def __init__(self, size): + self._size = size + + def size(self): + return self._size + + def rank(self): + return 0 + + +def _load_named(rel: str, name: str, extra_ns=None, class_name=None): + src = (ROOT / rel).read_text() + tree = ast.parse(src) + body = tree.body + if class_name: + body = next( + node.body + for node in tree.body + if isinstance(node, ast.ClassDef) and node.name == class_name + ) + target = next( + node + for node in body + if isinstance(node, ast.ClassDef) + and node.name == name + or isinstance(node, ast.FunctionDef) + and node.name == name + ) + if isinstance(target, ast.FunctionDef): + target.decorator_list = [] + mod = ast.Module(body=[target], type_ignores=[]) + ast.fix_missing_locations(mod) + ns = { + "torch": torch, + "F": F, + "Optional": Optional, + "List": List, + "os": __import__("os"), + "warnings": __import__("warnings"), + "_use_accuracy_compatible": _use_accuracy_compatible, + "get_tensor_model_parallel_group_if_none": lambda g: g, + "LinearWithGradAccumulationAndAsyncCommunication": _SentinelApply, + "parallel_state": SimpleNamespace(get_tensor_model_parallel_world_size=lambda: _TP["size"]), + "custom_backward": _custom_backward, + "Variable": torch.autograd.Variable, + } + if extra_ns: + ns.update(extra_ns) + exec(compile(mod, rel, "exec"), ns) + return ns[name] + + +_EmbedFp32MainGrad = _load_named("megatron/core/tensor_parallel/layers.py", "_EmbedFp32MainGrad") +linear_with_grad_accumulation_and_async_allreduce = _load_named( + "megatron/core/tensor_parallel/layers.py", "linear_with_grad_accumulation_and_async_allreduce" +) +linear_with_grad_accumulation_and_async_allreduce.warned = True +deallocate_output_tensor = _load_named( + "megatron/core/pipeline_parallel/schedules.py", "deallocate_output_tensor" +) +backward_step = _load_named("megatron/core/pipeline_parallel/schedules.py", "backward_step") + + +def _ref_embed_fp32_wgrad(weight_bf16, ids, grad_out): + table = weight_bf16.detach().clone().requires_grad_(True) + looked = F.embedding(ids, table) + (gw,) = torch.autograd.grad(looked, table, grad_outputs=grad_out) + return gw.float() + + +def _cuda_bf16(values, shape, device): + return torch.tensor(values, device=device, dtype=torch.bfloat16).reshape(shape) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestEmbedFp32MainGradCuda(unittest.TestCase): + def test_embedding_configuration_controls_gradient_destination(self): + forward = _load_named( + "megatron/core/tensor_parallel/layers.py", + "forward", + {"_EmbedFp32MainGrad": _EmbedFp32MainGrad}, + class_name="VocabParallelEmbedding", + ) + ids = torch.tensor([1, 1, 3], device="cuda") + instances = [] + for enabled in (False, True): + weight = torch.ones(4, 8, device="cuda", dtype=torch.bfloat16, requires_grad=True) + weight.main_grad = torch.zeros_like(weight, dtype=torch.float32) + instances.append( + SimpleNamespace( + config=SimpleNamespace(dsa_accuracy_compatible=enabled), + deterministic_mode=True, + tp_group=_FakeGroup(1), + weight=weight, + reduce_scatter_embeddings=False, + ) + ) + for instance in instances: + enabled = instance.config.dsa_accuracy_compatible + with patch.dict( + os.environ, + { + "MODEL_REPRO_TWO_FP32_ACCUM": str(int(not enabled)), + "USE_ACCURACY_COMPATIBLE": str(int(not enabled)), + }, + ): + output = forward(instance, ids) + output.sum().backward() + expected = torch.zeros(4, 8, device="cuda", dtype=torch.float32) + expected[1] = 2 + expected[3] = 1 + if enabled: + self.assertIsNone(instance.weight.grad) + torch.testing.assert_close(instance.weight.main_grad, expected, atol=0, rtol=0) + else: + torch.testing.assert_close(instance.weight.grad.float(), expected, atol=0, rtol=0) + self.assertEqual(torch.count_nonzero(instance.weight.main_grad).item(), 0) + + def test_repeated_indices_matches_independent_full_table_autograd(self): + device = torch.device("cuda") + vocab, dim = 8, 4 + weight = _cuda_bf16( + [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 10, + 11, + 12, + 13, + 14, + 15, + 16, + 17, + 18, + 19, + 20, + 21, + 22, + 23, + 24, + 25, + 26, + 27, + 28, + 29, + 30, + 31, + 32, + ], + (vocab, dim), + device, + ) + ids = torch.tensor([[1, 3, 1, 5], [3, 3, 0, 1]], device=device) + grad_out = _cuda_bf16(list(range(1, 33)), (2, 4, dim), device) + w = weight.clone().requires_grad_(True) + w.main_grad = torch.zeros(vocab, dim, device=device, dtype=torch.float32) + w.grad_added_to_main_grad = False + out = _EmbedFp32MainGrad.apply(w, ids) + self.assertEqual(out.dtype, torch.bfloat16) + torch.testing.assert_close(out, w[ids], atol=0, rtol=0) + out.backward(grad_out) + ref = _ref_embed_fp32_wgrad(weight, ids, grad_out) + torch.testing.assert_close(w.main_grad, ref, atol=0, rtol=0) + self.assertIsNone(w.grad) + self.assertTrue(w.grad_added_to_main_grad) + unused = [i for i in range(vocab) if i not in set(ids.reshape(-1).tolist())] + self.assertTrue((w.main_grad[unused] == 0).all()) + + def test_main_grad_accumulates_two_backwards_unused_rows_stay_zero(self): + device = torch.device("cuda") + vocab, dim = 6, 3 + weight = _cuda_bf16(list(range(1, 19)), (vocab, dim), device) + ids_a = torch.tensor([2, 2, 4], device=device) + ids_b = torch.tensor([4, 1, 2], device=device) + go_a = _cuda_bf16([1, 2, 3, 4, 5, 6, 7, 8, 9], (3, dim), device) + go_b = _cuda_bf16([2, 1, 0, 1, 2, 3, 4, 5, 6], (3, dim), device) + w = weight.clone().requires_grad_(True) + w.main_grad = torch.zeros(vocab, dim, device=device, dtype=torch.float32) + w.grad_added_to_main_grad = False + _EmbedFp32MainGrad.apply(w, ids_a).backward(go_a) + first = w.main_grad.clone() + self.assertTrue(w.grad_added_to_main_grad) + _EmbedFp32MainGrad.apply(w, ids_b).backward(go_b) + ref = _ref_embed_fp32_wgrad(weight, ids_a, go_a) + _ref_embed_fp32_wgrad( + weight, ids_b, go_b + ) + torch.testing.assert_close(w.main_grad, ref, atol=0, rtol=0) + self.assertFalse(torch.equal(first, w.main_grad)) + self.assertTrue((w.main_grad[0] == 0).all()) + self.assertTrue((w.main_grad[3] == 0).all()) + self.assertTrue((w.main_grad[5] == 0).all()) + self.assertIsNone(w.grad) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestLinearTp1Native(unittest.TestCase): + def setUp(self): + _UAC["on"] = True + _SentinelApply.last = None + + def tearDown(self): + _UAC["on"] = False + + def test_tp1_forward_dgrad_wgrad_bias_matches_f_linear(self): + device = torch.device("cuda") + x = _cuda_bf16( + [1, 2, 0, -1, 1, 0, 2, 1, -2, 0, 1, 1, 2, 0, 1], (5, 3), device + ).requires_grad_(True) + w = _cuda_bf16([1, 0, -1, 0, 1, 1, 1, -1, 0, 0, 1, -1], (4, 3), device).requires_grad_(True) + b = _cuda_bf16([1, -1, 0, 2], (4,), device).requires_grad_(True) + out = linear_with_grad_accumulation_and_async_allreduce( + x, w, b, False, False, False, None, 0, _FakeGroup(1), dsa_accuracy_compatible=True + ) + xref = x.detach().clone().requires_grad_(True) + wref = w.detach().clone().requires_grad_(True) + bref = b.detach().clone().requires_grad_(True) + ref = F.linear(xref, wref, bref) + torch.testing.assert_close(out, ref, atol=0, rtol=0) + self.assertIsNone(_SentinelApply.last) + go = _cuda_bf16( + [1, 1, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 0, 1, 1], (5, 4), device + ) + out.backward(go) + F.linear(xref, wref, bref).backward(go) + torch.testing.assert_close(x.grad, xref.grad, atol=0, rtol=0) + torch.testing.assert_close(w.grad, wref.grad, atol=0, rtol=0) + torch.testing.assert_close(b.grad, bref.grad, atol=0, rtol=0) + + def test_global_accuracy_without_dsa_keeps_native_custom_function(self): + _UAC["on"] = True + device = torch.device("cuda") + x = torch.ones(2, 3, device=device, dtype=torch.bfloat16, requires_grad=True) + w = torch.ones(4, 3, device=device, dtype=torch.bfloat16) + linear_with_grad_accumulation_and_async_allreduce( + x, w, None, False, False, False, None, 0, _FakeGroup(1) + ) + self.assertIsNotNone(_SentinelApply.last) + self.assertTrue(torch.equal(_SentinelApply.last[0], x)) + + def test_tp2_or_allreduce_skips_tp1_matmul_path(self): + device = torch.device("cuda") + x = torch.ones(2, 3, device=device, dtype=torch.bfloat16, requires_grad=True) + w = torch.ones(4, 3, device=device, dtype=torch.bfloat16) + linear_with_grad_accumulation_and_async_allreduce( + x, w, None, False, False, False, None, 0, _FakeGroup(2), dsa_accuracy_compatible=True + ) + self.assertIsNotNone(_SentinelApply.last) + _SentinelApply.last = None + linear_with_grad_accumulation_and_async_allreduce( + x, w, None, False, True, False, None, 0, _FakeGroup(1), dsa_accuracy_compatible=True + ) + self.assertIsNotNone(_SentinelApply.last) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestPipelineHelpersTp1(unittest.TestCase): + def tearDown(self): + _UAC["on"] = False + _TP["size"] = 1 + _CUSTOM_BWD["calls"] = [] + + def test_deallocate_uac_tp1_preserves_storage(self): + _UAC["on"] = True + _TP["size"] = 1 + t = torch.arange(4.0, device="cuda", requires_grad=True) + data_before = t.data.clone() + deallocate_output_tensor( + t, True, SimpleNamespace(dsa_accuracy_compatible=True, tensor_model_parallel_size=1) + ) + torch.testing.assert_close(t.data, data_before, atol=0, rtol=0) + self.assertEqual(tuple(t.shape), (4,)) + + def test_deallocate_off_or_tp2_still_frees(self): + _UAC["on"] = False + _TP["size"] = 1 + t = torch.arange(4.0, device="cuda") + deallocate_output_tensor(t, True) + self.assertEqual(tuple(t.shape), (1,)) + _UAC["on"] = True + _TP["size"] = 2 + t2 = torch.arange(4.0, device="cuda") + deallocate_output_tensor(t2, True) + self.assertEqual(tuple(t2.shape), (1,)) + + def test_backward_step_uac_tp1_uses_autograd_not_custom(self): + _UAC["on"] = True + _CUSTOM_BWD["calls"] = [] + x = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], device="cuda", requires_grad=True) + y = x * 3 + go = torch.ones_like(y) + cfg = SimpleNamespace( + timers=None, + grad_scale_func=None, + deallocate_pipeline_outputs=True, + tensor_model_parallel_size=1, + dsa_accuracy_compatible=True, + ) + gin = backward_step(x, y, go, cfg) + torch.testing.assert_close(gin, go * 3, atol=0, rtol=0) + self.assertEqual(_CUSTOM_BWD["calls"], []) + + def test_backward_step_off_or_tp2_uses_custom_backward(self): + x = torch.tensor([[1.0, 2.0], [3.0, 4.0]], device="cuda", requires_grad=True) + y = x * 2 + go = torch.ones_like(y) + cfg = SimpleNamespace( + timers=None, + grad_scale_func=None, + deallocate_pipeline_outputs=True, + tensor_model_parallel_size=1, + dsa_accuracy_compatible=False, + ) + _UAC["on"] = False + _CUSTOM_BWD["calls"] = [] + gin = backward_step(x, y, go, cfg) + self.assertEqual(len(_CUSTOM_BWD["calls"]), 1) + torch.testing.assert_close(gin, go * 2, atol=0, rtol=0) + + x2 = torch.tensor([[1.0, 2.0], [3.0, 4.0]], device="cuda", requires_grad=True) + y2 = x2 * 2 + go2 = torch.ones_like(y2) + cfg.dsa_accuracy_compatible = True + cfg.tensor_model_parallel_size = 2 + _UAC["on"] = True + _CUSTOM_BWD["calls"] = [] + gin2 = backward_step(x2, y2, go2, cfg) + self.assertEqual(len(_CUSTOM_BWD["calls"]), 1) + torch.testing.assert_close(gin2, go2 * 2, atol=0, rtol=0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit_tests/tensor_parallel/test_layers.py b/tests/unit_tests/tensor_parallel/test_layers.py index dbc27f502c6..be5795b2642 100644 --- a/tests/unit_tests/tensor_parallel/test_layers.py +++ b/tests/unit_tests/tensor_parallel/test_layers.py @@ -1,12 +1,42 @@ # Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +from types import SimpleNamespace + import pytest import torch -from megatron.core.tensor_parallel.layers import linear_with_frozen_weight +from megatron.core.tensor_parallel.layers import ( + _expert_grads_need_own_dp_domain, + linear_with_frozen_weight, +) from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region from tests.unit_tests.test_utilities import Utils +def test_expert_grads_need_own_dp_domain_etp_lt_tp(): + """EP=1 / ETP=1 / TP=2 expert wgrad must leave the dense dp_cp bucket.""" + frozen = SimpleNamespace( + expert_model_parallel_size=1, + tensor_model_parallel_size=2, + expert_tensor_parallel_size=1, + dsa_accuracy_compatible=True, + ) + assert _expert_grads_need_own_dp_domain(frozen) is True + frozen.dsa_accuracy_compatible = False + assert _expert_grads_need_own_dp_domain(frozen) is False + eq = SimpleNamespace( + expert_model_parallel_size=1, tensor_model_parallel_size=2, expert_tensor_parallel_size=2 + ) + assert _expert_grads_need_own_dp_domain(eq) is False + ep2 = SimpleNamespace( + expert_model_parallel_size=2, tensor_model_parallel_size=2, expert_tensor_parallel_size=1 + ) + assert _expert_grads_need_own_dp_domain(ep2) is True + missing = SimpleNamespace( + expert_model_parallel_size=1, tensor_model_parallel_size=2, expert_tensor_parallel_size=None + ) + assert _expert_grads_need_own_dp_domain(missing) is False + + @pytest.mark.parametrize("tensor_parallel,allreduce_dgrad", [(1, False), (8, True)]) def test_LinearWithFrozenWeight(tensor_parallel, allreduce_dgrad): Utils.initialize_model_parallel(tensor_parallel, 1) diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py index 642aeeb126f..466eefe0b73 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py @@ -27,6 +27,7 @@ DSAttention, DSAttentionSubmodules, FusedDSAIndexerLoss, + _AccuracyCompatibleSoftmax, _run_sparse_attention, _validate_nonpacked_cp_uniform_length, compute_dsa_indexer_loss, @@ -68,6 +69,59 @@ def mock_hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor return x * scale +class TestAccuracyCompatibleDSA: + """Test the opt-in full-score DSA alignment path.""" + + def test_explicit_softmax_backward_matches_formula(self): + logits = torch.randn(2, 3, 5, device="cuda", requires_grad=True) + valid_mask = torch.ones_like(logits, dtype=torch.bool) + valid_mask[..., -1] = False + grad_output = torch.randn_like(logits) + + probabilities = _AccuracyCompatibleSoftmax.apply(logits, valid_mask) + probabilities.backward(grad_output) + expected = probabilities.detach() * ( + grad_output - (grad_output * probabilities.detach()).sum(dim=-1, keepdim=True) + ) + expected = expected.masked_fill(~valid_mask, 0.0) + + assert torch.equal(logits.grad, expected) + assert torch.equal(probabilities[..., -1], torch.zeros_like(probabilities[..., -1])) + + def test_accuracy_compatible_switch_defaults_off(self, monkeypatch): + query = torch.randn(8, 1, 2, 8, device="cuda", dtype=torch.bfloat16) + key = torch.randn_like(query) + value = torch.randn(8, 1, 2, 4, device="cuda", dtype=torch.bfloat16) + indices = torch.arange(8, device="cuda").view(1, 8, 1) + original = unfused_dsa_fn + calls = [] + + def capture(*args, **kwargs): + calls.append(kwargs.get("accuracy_compatible")) + return original(*args, **kwargs) + + monkeypatch.setattr( + "megatron.core.transformer.experimental_attention_variant.dsa.unfused_dsa_fn", capture + ) + common = dict( + absorbed_mla=False, + query=query, + key=key, + value=value, + up_v_weight=None, + topk_indices=indices, + softmax_scale=query.size(-1) ** -0.5, + mask=None, + varlen_starts=None, + varlen_ends=None, + key_positions=None, + ) + _run_sparse_attention(config=SimpleNamespace(), **common) + _run_sparse_attention(config=SimpleNamespace(dsa_accuracy_compatible=True), **common) + + assert calls == [False, True] + + class TestDSAIndexShareHelpers: """Test cross-layer top-k sharing helpers.""" diff --git a/tests/unit_tests/transformer/moe/test_accuracy_migration.py b/tests/unit_tests/transformer/moe/test_accuracy_migration.py new file mode 100644 index 00000000000..620f71f44ad --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_accuracy_migration.py @@ -0,0 +1,247 @@ +"""CUDA unit tests for migrated MoE accuracy-compatible production paths.""" + +from __future__ import annotations + +import ast +import unittest +from pathlib import Path +from types import SimpleNamespace +from typing import Optional + +import torch + +ROOT = Path(__file__).resolve().parents[4] +_UAC = {"on": False} + + +def _use_accuracy_compatible(): + return _UAC["on"] + + +class MoECudaGraphPartialCaptureSignal(Exception): + pass + + +class NVLSAllGatherVDispatcher: + pass + + +def _load_fn(rel: str, name: str): + src = (ROOT / rel).read_text() + tree = ast.parse(src) + class_name = ( + "MoEFlexTokenDispatcher" + if rel.endswith("token_dispatcher.py") + else "MoELayer" if rel.endswith("moe_layer.py") else None + ) + body = tree.body + if class_name: + body = next( + node.body + for node in tree.body + if isinstance(node, ast.ClassDef) and node.name == class_name + ) + target = next(node for node in body if isinstance(node, ast.FunctionDef) and node.name == name) + target.decorator_list = [] + mod = ast.Module(body=[target], type_ignores=[]) + ast.fix_missing_locations(mod) + ns = { + "torch": torch, + "Optional": Optional, + "_use_accuracy_compatible": _use_accuracy_compatible, + "MoECudaGraphPartialCaptureSignal": MoECudaGraphPartialCaptureSignal, + "NVLSAllGatherVDispatcher": NVLSAllGatherVDispatcher, + "Tuple": tuple, + } + exec(compile(mod, rel, "exec"), ns) + return ns[name] + + +token_dispatch = _load_fn("megatron/core/transformer/moe/token_dispatcher.py", "token_dispatch") +token_combine = _load_fn("megatron/core/transformer/moe/token_dispatcher.py", "token_combine") +moe_forward = _load_fn("megatron/core/transformer/moe/moe_layer.py", "forward") +moe_postprocess = _load_fn("megatron/core/transformer/moe/moe_layer.py", "postprocess") +topk_routing_with_score_function = _load_fn( + "megatron/core/transformer/moe/moe_utils.py", "topk_routing_with_score_function" +) + + +class _FakeComm: + def __init__(self): + self.calls = [] + self.dispatched_probs = None + + def dispatch(self, hidden, async_finish, allocate_on_comm_stream): + self.calls.append(("dispatch", async_finish, allocate_on_comm_stream)) + self.dispatched_probs = hidden + return hidden + + def combine(self, hidden, async_finish, allocate_on_comm_stream): + self.calls.append(("combine", async_finish, allocate_on_comm_stream)) + return hidden + + def combine_postprocess(self, output): + return output + + +class _Flex: + @property + def config(self): + return SimpleNamespace(dsa_accuracy_compatible=_UAC["on"]) + + def __init__(self): + self.shared_experts = None + self._comm_manager = _FakeComm() + + token_dispatch = token_dispatch + token_combine = token_combine + + +class _MoE: + @property + def config(self): + return SimpleNamespace( + dsa_accuracy_compatible=_UAC["on"], + sequence_parallel=True, + moe_shared_expert_overlap=False, + moe_latent_size=0, + fp8=False, + fp4=False, + ) + + def __init__(self): + self.training = False + self.attn_tp_group = SimpleNamespace(size=lambda: 1) + + self.token_dispatcher = SimpleNamespace(combine_postprocess=lambda x: x) + self.shared_expert_overlap = False + self.fwd_execution_map = {"route", "expert_compute", "postprocess"} + self.moe_layer_recompute = False + self._accuracy_shared_input = None + self.order = [] + + def shared_experts_compute(self, x): + self.order.append("shared") + return x * 2 + + def route(self, x, padding_mask): + self.order.append("route") + return x, x + + def preprocess(self, x, probs, routing_map): + self.order.append("preprocess") + return x, probs + + def dispatch(self, x, probs): + self.order.append("dispatch") + return x, probs + + def routed_experts_compute(self, x, probs): + self.order.append("routed") + return x * 3, None + + def combine(self, x): + self.order.append("combine") + return x + + postprocess = moe_postprocess + forward = moe_forward + + +class TestAccuracyMigration(unittest.TestCase): + @classmethod + def setUpClass(cls): + if not torch.cuda.is_available(): + raise RuntimeError("CUDA unavailable") + torch.cuda.set_device(0) + + def setUp(self): + torch.cuda.set_device(0) + if not torch.cuda.is_available(): + raise RuntimeError("CUDA unavailable") + + def test_token_dispatch_combine_flags(self): + flex = _Flex() + x = torch.ones(2, 4, device="cuda") + _UAC["on"] = True + flex.token_dispatch(x, async_finish=True, allocate_on_comm_stream=True) + flex.token_combine(x, async_finish=True, allocate_on_comm_stream=True) + self.assertEqual(flex._comm_manager.calls[0][1:], (False, False)) + self.assertEqual(flex._comm_manager.calls[1][1:], (False, False)) + flex._comm_manager.calls.clear() + _UAC["on"] = False + flex.token_dispatch(x, async_finish=True, allocate_on_comm_stream=False) + flex.token_combine(x, async_finish=False, allocate_on_comm_stream=True) + self.assertEqual(flex._comm_manager.calls[0][1:], (True, False)) + self.assertEqual(flex._comm_manager.calls[1][1:], (False, True)) + + def test_forward_shared_order_and_value(self): + layer = _MoE() + x = torch.ones(2, 3, 4, device="cuda", requires_grad=True) + _UAC["on"] = True + out, _ = layer.forward(x) + self.assertEqual( + layer.order, ["route", "preprocess", "dispatch", "routed", "combine", "shared"] + ) + self.assertTrue(torch.equal(out, 5 * x.detach())) + out.sum().backward() + self.assertTrue(torch.equal(x.grad, torch.full_like(x, 5))) + self.assertIsNone(getattr(layer, "_accuracy_shared_input", None)) + layer.order.clear() + x2 = torch.ones(2, 3, 4, device="cuda", requires_grad=True) + out2, _ = layer.forward(x2) + self.assertTrue(torch.equal(out2, 5 * x2.detach())) + self.assertIsNone(getattr(layer, "_accuracy_shared_input", None)) + layer = _MoE() + _UAC["on"] = False + x3 = torch.ones(2, 3, 4, device="cuda", requires_grad=True) + out3, _ = layer.forward(x3) + self.assertEqual( + layer.order, ["shared", "route", "preprocess", "dispatch", "routed", "combine"] + ) + self.assertTrue(torch.equal(out3, 5 * x3.detach())) + + def test_postprocess_mixed_dtype(self): + layer = _MoE() + routed = torch.ones(2, 4, device="cuda", dtype=torch.bfloat16, requires_grad=True) + shared = torch.ones(2, 4, device="cuda", dtype=torch.float32, requires_grad=True) + _UAC["on"] = True + out = layer.postprocess(routed, shared) + self.assertEqual(out.dtype, torch.bfloat16) + out.float().sum().backward() + self.assertIsNotNone(routed.grad) + self.assertIsNotNone(shared.grad) + routed2 = torch.ones(2, 4, device="cuda", dtype=torch.bfloat16, requires_grad=True) + shared2 = torch.ones(2, 4, device="cuda", dtype=torch.float32, requires_grad=True) + _UAC["on"] = False + out2 = layer.postprocess(routed2, shared2) + self.assertEqual(out2.dtype, torch.float32) + + def test_topk_routing_output_dtype_and_gradients(self): + _UAC["on"] = True + logits = torch.tensor( + [[1.0, 2.0, 0.5, 3.0], [0.2, 4.0, 1.5, 0.8]], device="cuda", dtype=torch.bfloat16 + ) + logits = logits.clone().requires_grad_(True) + for score in ("sigmoid", "sqrtsoftplus"): + probs, _idx = topk_routing_with_score_function( + logits, + topk=2, + score_function=score, + dense_output=True, + fused=False, + router_replay=None, + ) + self.assertEqual(probs.dtype, logits.dtype) + self.assertTrue( + torch.allclose( + probs.float().sum(dim=-1), torch.ones(probs.size(0), device="cuda"), atol=0.01 + ) + ) + probs.sum().backward() + self.assertTrue(torch.isfinite(logits.grad).all()) + logits.grad = None + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit_tests/transformer/moe/test_routers.py b/tests/unit_tests/transformer/moe/test_routers.py index 9f33dd01920..9f33079c9a9 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -62,6 +62,41 @@ def test_constructor(self): num_weights = sum([p.numel() for p in self.router.parameters()]) assert num_weights == 12 * 4, num_weights + @pytest.mark.internal + def test_router_accuracy_compatible_gating(self): + hidden_states = torch.randn( + (3, 1, self.router.config.hidden_size), device="cuda", dtype=torch.bfloat16 + ) + self.router.config.router_accuracy_compatible = True + + logits = self.router.gating(hidden_states) + expected = torch.mm( + hidden_states.reshape(-1, hidden_states.shape[-1]).float(), + self.router.weight.float().t(), + ).view(3, 1, -1) + + assert logits.dtype == torch.float32 + assert torch.equal(logits, expected) + + @pytest.mark.internal + def test_default_router_gating_stays_native(self, monkeypatch): + expected = torch.randn((3, 1, self.router.config.num_moe_experts)) + called = False + + def fake_router_gating_linear(inp, weight, bias, router_dtype): + nonlocal called + called = True + return expected + + monkeypatch.setattr( + "megatron.core.transformer.moe.router.router_gating_linear", fake_router_gating_linear + ) + hidden_states = torch.randn((3, 1, self.router.config.hidden_size), dtype=torch.bfloat16) + + assert self.router.config.router_accuracy_compatible is False + assert self.router.gating(hidden_states) is expected + assert called + @pytest.mark.internal @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.parametrize("moe_router_pre_softmax", [(True), (False)]) diff --git a/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py b/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py new file mode 100644 index 00000000000..ea0decdba51 --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py @@ -0,0 +1,77 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import pytest +import torch + +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.mlp import MLPSubmodules +from megatron.core.transformer.moe.experts import SequentialMLP +from megatron.core.transformer.transformer_config import TransformerConfig +from tests.unit_tests.test_utilities import Utils + + +class TestSequentialExpertPadding: + def setup_method(self): + Utils.initialize_model_parallel(tensor_model_parallel_size=1, expert_model_parallel_size=1) + model_parallel_cuda_manual_seed(123) + + def teardown_method(self): + Utils.destroy_model_parallel() + + @pytest.mark.parametrize("enabled,grouped", [(False, False), (True, False), (True, True)]) + @pytest.mark.parametrize("first_count", [0, 3, 16, 17]) + def test_expert_gemm_rows_preserve_storage_contract(self, enabled, grouped, first_count): + config = TransformerConfig( + num_layers=1, + hidden_size=32, + num_attention_heads=4, + ffn_hidden_size=64, + moe_ffn_hidden_size=64, + num_moe_experts=2, + moe_router_topk=1, + moe_router_pre_softmax=True, + add_bias_linear=False, + gated_linear_unit=True, + activation_func=torch.nn.functional.silu, + bias_activation_fusion=False, + params_dtype=torch.bfloat16, + use_accuracy_compatible=enabled, + dsa_accuracy_compatible=True, + moe_grouped_gemm=grouped, + ) + experts = SequentialMLP( + 2, + config, + MLPSubmodules(linear_fc1=ColumnParallelLinear, linear_fc2=RowParallelLinear), + pg_collection=ProcessGroupCollection.use_mpu_process_groups(), + ).cuda() + observed = [] + + def record_rows(module, inputs): + tokens, probs = inputs + observed.append((tokens.shape[0], probs.shape[0])) + + handles = [ + expert.register_forward_pre_hook(record_rows) for expert in experts.local_experts + ] + counts = torch.tensor([first_count, 19], dtype=torch.int64) + tokens = torch.randn( + first_count + 19, 32, device="cuda", dtype=torch.bfloat16, requires_grad=True + ) + probs = torch.ones(first_count + 19, device="cuda", dtype=torch.float32, requires_grad=True) + try: + output, bias = experts(tokens, counts, probs) + expected_rows = 32 if enabled and not grouped and 0 < first_count < 17 else first_count + assert observed == [(expected_rows, expected_rows), (19, 19)] + assert output.shape == tokens.shape + assert bias is None + output.float().sum().backward() + assert tokens.grad.shape == tokens.shape + assert probs.grad.shape == probs.shape + assert torch.isfinite(tokens.grad).all() + assert torch.isfinite(probs.grad).all() + finally: + for handle in handles: + handle.remove() diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index c3c3944e007..6f4c530b0c3 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -29,6 +29,7 @@ process_mtp_loss, roll_tensor, ) +from megatron.core.transformer.torch_norm import WrappedTorchNorm from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import get_batch_on_this_cp_rank, is_te_min_version, unwrap_model from megatron.training.argument_utils import gpt_config_from_args, hybrid_config_from_args @@ -82,6 +83,32 @@ def _create_config_and_mtp_block_spec(self, tp, cp, use_te=False): ) return config, mtp_block_spec + def test_accuracy_compatible_norms_override_te_mtp_norms(self): + """Accuracy mode routes all MTP-owned norms through native Torch RMSNorm.""" + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + config = TransformerConfig( + mtp_num_layers=1, + num_layers=1, + hidden_size=64, + num_attention_heads=8, + normalization="RMSNorm", + norm_accuracy_compatible=True, + use_cpu_initialization=True, + ) + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec() + mtp_block_spec = get_gpt_mtp_block_spec( + config=config, spec=transformer_layer_spec, use_transformer_engine=True + ) + mtp_layer_spec = mtp_block_spec.layer_specs[0] + + assert mtp_layer_spec.submodules.enorm is WrappedTorchNorm + assert mtp_layer_spec.submodules.hnorm is WrappedTorchNorm + assert mtp_layer_spec.submodules.layer_norm is WrappedTorchNorm + final_norm = mtp_layer_spec.submodules.layer_norm( + config=config, hidden_size=config.hidden_size, eps=config.layernorm_epsilon + ) + assert isinstance(final_norm, torch.nn.RMSNorm) + def test_mtp_detach_heads_config(self): """Test that mtp_detach_heads config defaults to False.""" config = TransformerConfig( diff --git a/tests/unit_tests/transformer/test_torch_norm.py b/tests/unit_tests/transformer/test_torch_norm.py new file mode 100644 index 00000000000..34d37ae4662 --- /dev/null +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -0,0 +1,40 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +import pytest +import torch + +from megatron.core.transformer.torch_norm import WrappedTorchNorm +from megatron.core.transformer.transformer_config import TransformerConfig + + +def _config(**overrides): + values = { + "num_layers": 1, + "hidden_size": 64, + "num_attention_heads": 4, + "normalization": "RMSNorm", + } + values.update(overrides) + return TransformerConfig(**values) + + +def test_rmsnorm_uses_native_torch_implementation(): + config = _config(norm_accuracy_compatible=True, params_dtype=torch.bfloat16) + norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) + + assert isinstance(norm, torch.nn.RMSNorm) + assert norm.weight.dtype == torch.bfloat16 + + +def test_sequence_parallel_is_opt_in_for_native_norm(): + with pytest.raises(AssertionError, match="sequence parallel"): + WrappedTorchNorm( + config=_config(sequence_parallel=True, tensor_model_parallel_size=2), + hidden_size=64, + eps=1e-5, + ) + config = _config( + sequence_parallel=True, tensor_model_parallel_size=2, norm_accuracy_compatible=True + ) + norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) + assert all(parameter.sequence_parallel for parameter in norm.parameters())