diff --git a/megatron/core/distributed/finalize_model_grads.py b/megatron/core/distributed/finalize_model_grads.py index d5b1e77066f..8cbda1ca0fd 100644 --- a/megatron/core/distributed/finalize_model_grads.py +++ b/megatron/core/distributed/finalize_model_grads.py @@ -467,9 +467,7 @@ 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 not config.dsa_accuracy_compatible and num_tokens is not None - ) + loss_normalized_in_graph = _use_accuracy_compatible() and (num_tokens is not None) if loss_normalized_in_graph: num_tokens = None @@ -575,7 +573,6 @@ 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 103c2a9a7d6..b7de1013695 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -916,12 +916,8 @@ def __init__( for param in self.parameters(): if is_expert: - # 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)) + # Reduce the gradient on the expert_data_parallel group for expert linear layers + setattr(param, "allreduce", not self.expert_parallel) else: # Reduce the gradient on DP group setattr(param, "allreduce", True) @@ -2039,15 +2035,7 @@ def __init__( ) for param in self.parameters(): - # 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)), - ) + setattr(param, "allreduce", not (is_expert and self.expert_parallel)) # 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 1e1c367d89c..88bb070e105 100644 --- a/megatron/core/model_parallel_config.py +++ b/megatron/core/model_parallel_config.py @@ -150,9 +150,6 @@ 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 e71db578f14..a270161ddd6 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -99,17 +99,6 @@ 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 615e2c45196..fb7b8ef405e 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() and not self.config.dsa_accuracy_compatible: + if _use_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 4b01ba561b6..a76fe6e3a23 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -20,7 +20,6 @@ ) 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, @@ -58,13 +57,6 @@ ########## -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: @@ -73,11 +65,12 @@ 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=_get_standalone_norm(config, backend), + out_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False), out_proj=backend.row_parallel_linear(), ), metainfo={"fuse_input_layernorm": True}, @@ -109,10 +102,12 @@ 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 = ( - _get_standalone_norm(config, backend, for_qk=True) if config.qk_layernorm else IdentityOp + backend.layer_norm(rms_norm=rms_norm, for_qk=True) if config.qk_layernorm else IdentityOp ) attention = ModuleSpec( @@ -233,6 +228,7 @@ 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 = ( @@ -249,10 +245,12 @@ def get_transformer_layer_with_experimental_attention_variant_spec( input_layernorm = ( IdentityOp if attention.metainfo["fuse_input_layernorm"] - else _get_standalone_norm(config, backend) + else backend.layer_norm(rms_norm=rms_norm, for_qk=False) ) pre_mlp_layernorm = ( - IdentityOp if fuse_pre_mlp_layernorm else _get_standalone_norm(config, backend) + IdentityOp + if fuse_pre_mlp_layernorm + else backend.layer_norm(rms_norm=rms_norm, for_qk=False) ) layer_specs.append( @@ -319,8 +317,9 @@ 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=_get_standalone_norm(config, backend) + layer_specs=layer_specs, layer_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False) ) 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 63a1aa51b02..984840b3a87 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, config=config + mtp_model_layer_spec=transformer_layer_spec, backend=backend ) 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 c7406191dc2..32a61cf7efc 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -556,9 +556,6 @@ 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 a11f5c8cc3b..24f9a032c47 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -247,9 +247,6 @@ 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 f576bc05a58..e67c498e2cc 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, config=None): +def deallocate_output_tensor(out, deallocate_pipeline_outputs=False): '''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,23 +177,17 @@ def deallocate_output_tensor(out, deallocate_pipeline_outputs=False, config=None ''' 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, config) + deallocate_output_tensor(value, deallocate_pipeline_outputs) return # Handle list format if isinstance(out, list): for item in out: - deallocate_output_tensor(item, deallocate_pipeline_outputs, config) + deallocate_output_tensor(item, deallocate_pipeline_outputs) return # Base case: deallocate tensor @@ -574,10 +568,7 @@ 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: - _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 - ): + if config.deallocate_pipeline_outputs: custom_backward(output_tensor[0], output_tensor_grad[0]) else: torch.autograd.backward(output_tensor[0], grad_tensors=output_tensor_grad[0]) @@ -650,10 +641,7 @@ 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: - _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 - ): + if config.deallocate_pipeline_outputs: custom_backward(output_tensor_module, output_tensor_grad_module) else: torch.autograd.backward( @@ -1649,7 +1637,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, config) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) else: if not is_pp_first_stage(p2p_communicator.pp_group): # Send only since recv prefetched. @@ -1679,7 +1667,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, config) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) if recv_prev: input_tensors[next_forward_model_chunk_id].append( fwd_recv_buffer[k % fwd_recv_buffer_size] @@ -1757,7 +1745,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, config) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) # Async forward send / receive def pp_post_forward(output_tensor, vp_stage=None): @@ -1932,7 +1920,7 @@ def pp_post_backward(input_tensor_grad, vp_stage=None): tensor_shape=tensor_shape, ) ) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) # Put input_tensor and output_tensor_grad in data structures in the # right location. if recv_prev: @@ -1940,7 +1928,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, config) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) nvtx_range_pop(suffix="steady") # Run cooldown backward passes (flush out pipeline) for the last model chunk. @@ -2365,7 +2353,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, config) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) # Before running 1F1B, need to receive first forward tensor. # If all microbatches are run in warmup / cooldown phase, then no need to @@ -2420,7 +2408,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, config) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) # 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 b613c581ed2..2f927d5218d 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -59,51 +59,6 @@ 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, @@ -282,11 +237,7 @@ 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() - or config.use_accuracy_compatible - ) + self.deterministic_mode = config.deterministic_mode or _use_accuracy_compatible() self.config = config self.use_inference_optimized_reduce_scatter = ( @@ -348,11 +299,7 @@ def forward(self, input_): masked_input = input_ # Get the embeddings. if self.deterministic_mode: - _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] + output_parallel = self.weight[masked_input] else: # F.embedding currently has a non-deterministic backward function output_parallel = F.embedding(masked_input, self.weight) @@ -728,7 +675,6 @@ 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. @@ -794,12 +740,6 @@ 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, @@ -837,23 +777,6 @@ 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. @@ -998,11 +921,7 @@ def __init__( tensor=self.weight, is_parallel=True, dim=0, stride=stride ) - setattr( - self.weight, - "allreduce", - not (self.is_expert and _expert_grads_need_own_dp_domain(config)), - ) + setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) else: self.weight = None @@ -1024,11 +943,7 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr( - self.bias, - "allreduce", - not (self.is_expert and _expert_grads_need_own_dp_domain(config)), - ) + setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) else: self.register_parameter("bias", None) @@ -1076,13 +991,7 @@ 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, - dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), - **kwargs, - ) + return linear_with_grad_accumulation_and_async_allreduce(input, weight, *args, **kwargs) def forward( self, @@ -1130,10 +1039,6 @@ 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) @@ -1179,10 +1084,7 @@ def forward( if runtime_gather_output is not None: gather_output = runtime_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) - ): + if gather_output: # 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. @@ -1369,11 +1271,7 @@ 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 _expert_grads_need_own_dp_domain(config)), - ) + setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) if bias: if config.use_cpu_initialization: @@ -1391,11 +1289,7 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr( - self.bias, - "allreduce", - not (self.is_expert and _expert_grads_need_own_dp_domain(config)), - ) + setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) setattr(self.bias, "sequence_parallel", self.sequence_parallel) else: self.register_parameter("bias", None) @@ -1411,13 +1305,7 @@ 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, - dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), - **kwargs, - ) + return linear_with_grad_accumulation_and_async_allreduce(input, weight, *args, **kwargs) def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Forward of RowParallelLinear @@ -1467,10 +1355,6 @@ 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 8e0c135d725..dde238635c2 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -64,7 +64,6 @@ 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() @@ -100,15 +99,10 @@ def _unfused_absorbed_dsa_fn( ) attention_scores = attention_scores + index_mask.unsqueeze(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 - ) + 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 + ) # 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] @@ -116,25 +110,6 @@ 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, @@ -152,7 +127,6 @@ 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: @@ -169,7 +143,7 @@ def _run_sparse_attention( "Received absorbed layout with explicit value tensor." ) output = None - if not accuracy_compatible and dsa_kernels.use_fused_dsa_kernels(config): + if dsa_kernels.use_fused_dsa_kernels(config): output = dsa_kernels.run_fused_absorbed_sparse_attention( config, query, @@ -192,7 +166,6 @@ 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() @@ -209,7 +182,6 @@ def _run_sparse_attention( varlen_starts=varlen_starts, varlen_ends=varlen_ends, key_positions=key_positions, - accuracy_compatible=accuracy_compatible, ) @@ -1439,7 +1411,6 @@ 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. @@ -1486,27 +1457,6 @@ 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 @@ -1827,11 +1777,8 @@ def forward( skv = key.size(0) # Detach x and qr to prevent gradients of indexer from flowing back to the main model. - _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() + 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 67fc3106187..229d02a00b6 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -1347,24 +1347,35 @@ def forward( output_local_list = [] - 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. + 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. 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, tokens.new_zeros(num_pad_tokens, tokens.shape[1])), dim=0 + ( + 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, ) - 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 798eb4cabd6..5eb8de35f09 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -576,11 +576,7 @@ def postprocess(self, output: torch.Tensor, shared_expert_output: Optional[torch output, _ = self.fc2_latent_proj(output) if shared_expert_output is not None: - 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 + output = output + shared_expert_output elif ( isinstance(self.token_dispatcher, NVLSAllGatherVDispatcher) and self._latent_shared_expert_output is not None @@ -664,11 +660,7 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): hidden_states_router = hidden_states hidden_states_dispatch = hidden_states - 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) + 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 @@ -703,12 +695,6 @@ 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 6da367bb659..afea9dcbf34 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -404,7 +404,6 @@ 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], @@ -525,11 +524,7 @@ 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 dsa_accuracy_compatible - and not (drop_and_pad and num_out_tokens is not None) - ): + if _use_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 a1ca6cdbe42..c9aedd860c6 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -104,12 +104,6 @@ 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 1b26b87e4e4..490e4d8cc6d 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -304,7 +304,6 @@ 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, ) ) @@ -652,7 +651,6 @@ 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 @@ -1379,7 +1377,6 @@ 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) @@ -1529,9 +1526,6 @@ 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 ) @@ -1591,9 +1585,6 @@ 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 5a91a7da037..b20514ce6a4 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, WrappedTorchNorm +from megatron.core.transformer.torch_norm import LayerNormBuilder 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,9 +577,7 @@ class MultiTokenPredictionLayerSubmodules: def get_mtp_layer_spec( - mtp_model_layer_spec: ModuleSpec, - use_transformer_engine: bool, - config: Optional[TransformerConfig] = None, + mtp_model_layer_spec: ModuleSpec, use_transformer_engine: bool ) -> ModuleSpec: """Get the MTP layer spec. @@ -589,14 +587,11 @@ 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, - config: Optional[TransformerConfig] = None, + mtp_model_layer_spec: ModuleSpec, backend: BackendSpecProvider ) -> ModuleSpec: """Get the MTP layer spec. @@ -604,11 +599,7 @@ 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 = ( - WrappedTorchNorm - if config is not None and config.norm_accuracy_compatible - else backend.layer_norm() - ) + layer_norm_impl = backend.layer_norm() mtp_layer_spec = ModuleSpec( module=MultiTokenPredictionLayer, submodules=MultiTokenPredictionLayerSubmodules( @@ -1115,11 +1106,7 @@ def _get_embeddings( if self.config.mtp_detach_heads: decoder_input = decoder_input.detach() - _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 - ) + 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 @@ -1134,17 +1121,10 @@ 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) - 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 - ) + decoder_input = make_viewless_tensor(inp=decoder_input, requires_grad=True, keep_graph=True) hidden_states = apply_module(self.hnorm)(hidden_states) - 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 - ) + 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) @@ -1155,7 +1135,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 ) - elif not (self.config.dsa_accuracy_compatible and _tp_size <= 1): + else: 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 a7f9db5e60a..5948ae600f9 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -47,9 +47,7 @@ def __new__( assert not config.persist_layer_norm, f"persist_layer_norm 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.sequence_parallel, f"sequence parallel not supported by torch LayerNorm" assert ( not config.memory_efficient_layer_norm @@ -68,14 +66,7 @@ def __new__( else: raise Exception("Only LayerNorm, RMSNorm and L2Norm are currently supported") - 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 + return norm_cls(normalized_shape=hidden_size, eps=eps) class L2Norm(torch.nn.Module, LayerNormInterface): diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index b3b977f8053..0415035ffbe 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -590,11 +590,7 @@ 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. - _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 - ) + 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() @@ -698,11 +694,9 @@ 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. - _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 - ) + 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 34c02443365..bbcf413baee 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -186,16 +186,6 @@ 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"]}} ) @@ -278,13 +268,6 @@ 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.""" @@ -335,14 +318,6 @@ 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 5a584e59f42..904912c18d8 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -941,12 +941,9 @@ 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. - 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 - ) + 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 762e1fdd81d..ee535c29baf 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -1,4 +1,5 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. +import inspect import os import pytest @@ -20,54 +21,6 @@ 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 13cc7419c45..0a454b5d7ff 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,7 +65,6 @@ 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, @@ -370,23 +369,6 @@ 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 deleted file mode 100644 index 4564ce0adf8..00000000000 --- a/tests/unit_tests/models/test_local_spec_provider_linear.py +++ /dev/null @@ -1,15 +0,0 @@ -# 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 deleted file mode 100644 index 979fff5282f..00000000000 --- a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py +++ /dev/null @@ -1,384 +0,0 @@ -"""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 be5795b2642..dbc27f502c6 100644 --- a/tests/unit_tests/tensor_parallel/test_layers.py +++ b/tests/unit_tests/tensor_parallel/test_layers.py @@ -1,42 +1,12 @@ # Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. -from types import SimpleNamespace - import pytest import torch -from megatron.core.tensor_parallel.layers import ( - _expert_grads_need_own_dp_domain, - linear_with_frozen_weight, -) +from megatron.core.tensor_parallel.layers import 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 466eefe0b73..642aeeb126f 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,7 +27,6 @@ DSAttention, DSAttentionSubmodules, FusedDSAIndexerLoss, - _AccuracyCompatibleSoftmax, _run_sparse_attention, _validate_nonpacked_cp_uniform_length, compute_dsa_indexer_loss, @@ -69,59 +68,6 @@ 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 deleted file mode 100644 index 620f71f44ad..00000000000 --- a/tests/unit_tests/transformer/moe/test_accuracy_migration.py +++ /dev/null @@ -1,247 +0,0 @@ -"""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 9f33079c9a9..9f33dd01920 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -62,41 +62,6 @@ 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 deleted file mode 100644 index ea0decdba51..00000000000 --- a/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py +++ /dev/null @@ -1,77 +0,0 @@ -# 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 6f4c530b0c3..c3c3944e007 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -29,7 +29,6 @@ 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 @@ -83,32 +82,6 @@ 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 deleted file mode 100644 index 34d37ae4662..00000000000 --- a/tests/unit_tests/transformer/test_torch_norm.py +++ /dev/null @@ -1,40 +0,0 @@ -# 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())