From faaec38949e14f4e4d6e90b906a98a2f7334198f Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Wed, 16 Sep 2026 17:56:51 +0800 Subject: [PATCH] Revert "Use one accuracy-compatible mode for model and optimizer numerics (#17)" This reverts commit fd8419182d409f2ef704d664d6be4aa0b0ece47d. --- .../core/distributed/finalize_model_grads.py | 4 +- .../common/language_module/language_module.py | 4 +- ...rimental_attention_variant_module_specs.py | 2 +- megatron/core/optimizer/__init__.py | 23 ++- megatron/core/optimizer/clip_grads.py | 45 +----- megatron/core/optimizer/optimizer.py | 38 +---- megatron/core/optimizer/optimizer_config.py | 6 +- megatron/core/optimizer/reproducible_norm.py | 135 ---------------- megatron/core/pipeline_parallel/schedules.py | 12 +- megatron/core/tensor_parallel/layers.py | 18 +-- .../experimental_attention_variant/dsa.py | 6 +- megatron/core/transformer/moe/experts.py | 2 +- megatron/core/transformer/moe/moe_layer.py | 6 +- megatron/core/transformer/moe/moe_utils.py | 4 +- megatron/core/transformer/moe/router.py | 2 +- .../core/transformer/moe/token_dispatcher.py | 10 +- .../transformer/multi_token_prediction.py | 10 +- megatron/core/transformer/torch_norm.py | 4 +- .../core/transformer/transformer_block.py | 4 +- .../core/transformer/transformer_config.py | 23 ++- .../core/transformer/transformer_layer.py | 2 +- .../distributed/test_finalize_model_grads.py | 3 +- ...rimental_attention_variant_module_specs.py | 4 +- .../optimizer/test_reproducible_norm.py | 144 ------------------ .../test_accuracy_tp1_migration.py | 18 +-- .../unit_tests/tensor_parallel/test_layers.py | 4 +- .../test_attention_variant_dsa.py | 2 +- .../moe/test_accuracy_migration.py | 4 +- .../transformer/moe/test_routers.py | 5 +- .../moe/test_sequential_expert_padding.py | 2 +- .../test_multi_token_prediction.py | 3 +- .../unit_tests/transformer/test_torch_norm.py | 18 +-- 32 files changed, 120 insertions(+), 447 deletions(-) delete mode 100644 megatron/core/optimizer/reproducible_norm.py delete mode 100644 tests/unit_tests/optimizer/test_reproducible_norm.py diff --git a/megatron/core/distributed/finalize_model_grads.py b/megatron/core/distributed/finalize_model_grads.py index 530824e1464..d5b1e77066f 100644 --- a/megatron/core/distributed/finalize_model_grads.py +++ b/megatron/core/distributed/finalize_model_grads.py @@ -468,7 +468,7 @@ def finalize_model_grads( from ..transformer.module import _use_accuracy_compatible loss_normalized_in_graph = ( - _use_accuracy_compatible() and not config.uses_dsa_reference and num_tokens is not None + _use_accuracy_compatible() and not config.dsa_accuracy_compatible and num_tokens is not None ) if loss_normalized_in_graph: num_tokens = None @@ -575,7 +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.uses_dsa_reference: + 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/models/common/language_module/language_module.py b/megatron/core/models/common/language_module/language_module.py index 955346196ef..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() and not self.config.uses_dsa_reference: + 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] @@ -231,7 +231,7 @@ def compute_language_model_loss(self, labels: Tensor, logits: Tensor) -> Tensor: import hashlib as _hashlib _l = loss.detach().float().contiguous() - print( # pylint: disable=bad-builtin + print( f"\nper_token_loss: rank={torch.distributed.get_rank()} " f"shape={list(_l.shape)} " f"md5={_hashlib.md5(_l.cpu().numpy().tobytes()).hexdigest()}", 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 7efa02576a2..4b01ba561b6 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -60,7 +60,7 @@ def _get_standalone_norm(config: TransformerConfig, backend: BackendSpecProvider, *, for_qk=False): rms_norm = config.normalization == "RMSNorm" - if rms_norm and config.uses_dsa_reference: + if rms_norm and config.norm_accuracy_compatible: return WrappedTorchNorm return backend.layer_norm(rms_norm=rms_norm, for_qk=for_qk) diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 9d7e35e27e5..c7406191dc2 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -554,11 +554,11 @@ def _get_megatron_optimizer_based_on_param_groups( # set Adam class and weight decay mode depending # on source of optimizer (Torch or TE/Apex) - if config.use_accuracy_compatible and not config.use_precision_aware_optimizer: - adam_cls = torch.optim.AdamW - kwargs.update({"foreach": False, "fused": True}) - elif USING_PYTORCH_OPTIMIZER: + 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 @@ -587,7 +587,20 @@ def _get_megatron_optimizer_based_on_param_groups( if is_te_min_version("2.1.0.dev0"): kwargs.update({"store_param_remainders": config.store_param_remainders}) - optimizer = adam_cls(**kwargs) + # [对齐修复] use_accuracy_compatible=1 时强制 torch.optim.AdamW(fused=True), 与 + # PaddleFleet `paddle.optimizer.AdamW` 公式逐位对齐, 绕开 TE.FusedAdam / apex.FusedAdam + # 在 bias correction 顺序 / weight decay 时机 / eps 位置上的差异。 + # 仅在非 precision_aware_optimizer 分支启用, 避免触碰 TE 特有 kwargs。 + from ..transformer.module import _use_accuracy_compatible + if _use_accuracy_compatible() and not config.use_precision_aware_optimizer: + # torch.optim.AdamW 无 adam_w_mode 形参(AdamW 即 decoupled weight + # decay);若上游 TE/Apex 分支已注入该 kwarg,这里必须剔除,否则 + # torch.optim.AdamW 会抛 unexpected keyword argument。 + kwargs.pop("adam_w_mode", None) + kwargs["fused"] = True + optimizer = torch.optim.AdamW(**kwargs) + else: + optimizer = adam_cls(**kwargs) def init_state_fn(opt, config=None): for group in opt.param_groups: diff --git a/megatron/core/optimizer/clip_grads.py b/megatron/core/optimizer/clip_grads.py index 6345991e057..3c5491d39a1 100644 --- a/megatron/core/optimizer/clip_grads.py +++ b/megatron/core/optimizer/clip_grads.py @@ -50,37 +50,12 @@ from ..tensor_parallel import param_is_not_tensor_parallel_duplicate from ..transformer.module import param_is_not_shared from ..utils import get_data_parallel_group_if_dtensor, to_local_if_dtensor -from .reproducible_norm import ReproducibleL2Norm - - -@torch.no_grad() -def get_reproducible_grad_norm_bins( - grads_for_norm: List[torch.Tensor], - grad_stats_parallel_group: torch.distributed.ProcessGroup | None, -) -> torch.Tensor: - """Reduce exact FP32-square bins using the existing gradient ownership groups.""" - grads_for_norm = list(grads_for_norm) - data_parallel_group = None - for grad in grads_for_norm: - data_parallel_group = get_data_parallel_group_if_dtensor(grad, data_parallel_group) - grads_for_norm = [to_local_if_dtensor(grad) for grad in grads_for_norm] - accumulator = ReproducibleL2Norm(grads_for_norm[0].device if grads_for_norm else None) - bins = accumulator.zeros() - for grad in grads_for_norm: - if grad.layout != torch.strided: - raise TypeError("Reproducible clipping requires dense FP32 gradients") - bins = accumulator.accumulate(bins, grad) - if data_parallel_group: - torch.distributed.all_reduce(bins, group=data_parallel_group) - torch.distributed.all_reduce(bins, group=grad_stats_parallel_group) - return bins def get_grad_norm_fp32( grads_for_norm: Union[List[torch.Tensor], torch.Tensor], norm_type: Union[int, float] = 2, grad_stats_parallel_group: Optional[torch.distributed.ProcessGroup] = None, - use_accuracy_compatible: bool = False, ) -> float: """Calculate the p-norm of gradients in FP32 precision. @@ -105,12 +80,6 @@ def get_grad_norm_fp32( if isinstance(grads_for_norm, torch.Tensor): grads_for_norm = [grads_for_norm] - if use_accuracy_compatible: - if float(norm_type) != 2.0: - raise ValueError("Reproducible clipping supports only the L2 norm") - bins = get_reproducible_grad_norm_bins(grads_for_norm, grad_stats_parallel_group) - return ReproducibleL2Norm(bins.device).finish(bins)[0] - data_parallel_group = None for grad in grads_for_norm: data_parallel_group = get_data_parallel_group_if_dtensor(grad, data_parallel_group) @@ -215,14 +184,12 @@ def clip_grad_by_total_norm_fp32( dummy_overflow_buf = torch.zeros(1, dtype=torch.int, device='cuda') if isinstance(clip_coeff, torch.Tensor): clip_coeff.clamp_max_(1.0) - if multi_tensor_scale_tensor_impl is not None: - multi_tensor_applier( - multi_tensor_scale_tensor_impl, dummy_overflow_buf, [grads, grads], clip_coeff - ) - else: - multi_tensor_applier( - multi_tensor_scale_impl, dummy_overflow_buf, [grads, grads], clip_coeff.item() - ) + assert ( + multi_tensor_scale_tensor_impl is not None + ), "clip_coeff is tensor type. But multi_tensor_scale_tensor not available." + multi_tensor_applier( + multi_tensor_scale_tensor_impl, dummy_overflow_buf, [grads, grads], clip_coeff + ) elif clip_coeff < 1.0: multi_tensor_applier( multi_tensor_scale_impl, dummy_overflow_buf, [grads, grads], clip_coeff diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index e00c17b200f..4a74328d0d9 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -48,15 +48,9 @@ from ..dist_checkpointing.utils import add_prefix_for_sharding from ..transformer.module import param_is_not_shared from ..utils import log_single_rank -from .clip_grads import ( - clip_grad_by_total_norm_fp32, - count_zeros_fp32, - get_grad_norm_fp32, - get_reproducible_grad_norm_bins, -) +from .clip_grads import clip_grad_by_total_norm_fp32, count_zeros_fp32, get_grad_norm_fp32 from .grad_scaler import MegatronGradScaler from .optimizer_config import OptimizerConfig -from .reproducible_norm import ReproducibleL2Norm logger = getLogger(__name__) @@ -302,10 +296,7 @@ def get_grad_norm(self): """Compute and return grad norm.""" grads_for_norm = self.get_grads_for_grad_norm() total_norm = get_grad_norm_fp32( - grads_for_norm, - grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - use_accuracy_compatible=self.config.use_accuracy_compatible - and self.config.clip_grad > 0, + grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group() ) return total_norm @@ -317,10 +308,7 @@ def _compute_grad_norms_by_group(self) -> Dict[str, float]: if self.has_grad_norm_group(grad_norm_group): grouped_grads = self.get_grads_for_grad_norm(grad_norm_group) group_grad_norm = get_grad_norm_fp32( - grouped_grads, - grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - use_accuracy_compatible=self.config.use_accuracy_compatible - and self.config.clip_grad > 0, + grouped_grads, grad_stats_parallel_group=self.get_grad_stats_parallel_group() ) self.grad_norms_by_group[grad_norm_group] = group_grad_norm return self.grad_norms_by_group @@ -338,10 +326,7 @@ def clip_grad_norm(self, clip_grad: float) -> float: else: grads_for_norm = [] grad_norm = get_grad_norm_fp32( - grads_for_norm, - grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - use_accuracy_compatible=self.config.use_accuracy_compatible - and self.config.clip_grad > 0, + grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group() ) if clip_grad > 0.0 and params: @@ -1581,21 +1566,8 @@ def get_grad_stats_parallel_group(self) -> torch.distributed.ProcessGroup: ) return self.chained_optimizers[0].get_grad_stats_parallel_group() - @torch.no_grad() - def _get_reproducible_grad_norm(self, grad_norm_group=None): - bins = None - for optimizer in self.chained_optimizers: - part = get_reproducible_grad_norm_bins( - optimizer.get_grads_for_grad_norm(grad_norm_group), - optimizer.get_grad_stats_parallel_group(), - ) - bins = part if bins is None else bins + part - return ReproducibleL2Norm(bins.device).finish(bins)[0] - @torch.no_grad() def get_grad_norm(self): - if self.config.use_accuracy_compatible and self.config.clip_grad > 0: - return self._get_reproducible_grad_norm() if len(self.chained_optimizers) == 1: return self.chained_optimizers[0].get_grad_norm() if self.grads_states_parallel_group_is_shared(): @@ -1659,8 +1631,6 @@ def has_grad_norm_group(self, grad_norm_group: str) -> bool: def _get_grad_norm_for_group(self, grad_norm_group: str): """Compute gradient norm for a named parameter group.""" _validate_grad_norm_group(grad_norm_group) - if self.config.use_accuracy_compatible and self.config.clip_grad > 0: - return self._get_reproducible_grad_norm(grad_norm_group) if self.grads_states_parallel_group_is_shared(): grouped_grads = [] for optimizer in self.chained_optimizers: diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index ba3c8a6aa10..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. @@ -368,9 +371,6 @@ class OptimizerConfig: ################ # Miscellaneous ################ - use_accuracy_compatible: bool = False - """Select reference optimizer numerics and reproducible gradient clipping.""" - clip_grad: float = 1.0 """Gradient clipping based on global L2 norm.""" diff --git a/megatron/core/optimizer/reproducible_norm.py b/megatron/core/optimizer/reproducible_norm.py deleted file mode 100644 index 3380c3dbf84..00000000000 --- a/megatron/core/optimizer/reproducible_norm.py +++ /dev/null @@ -1,135 +0,0 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - -"""Layout- and partition-independent L2 norm of FP32 gradients. - -Squares are computed in FP32 on the gradient device. Their mantissas are -accumulated into base-65536 integer bins, which may be SUM-reduced across -owners before rounding the total once to FP32 and applying the native sqrt. -No gradient values or floating-point norm computation leave the device. - -The 20 limbs cover FP32 squares and the supported 2**40 global element count. -Each uncarried limb is bounded by 2**40 * 65535 < 2**56, so arbitrary parameter -and rank reduction orders cannot overflow int64. The last three bins count -infinities, NaNs, and elements. This deliberately costs more than a fused norm -and is intended only for explicitly enabled accuracy compatibility. -""" - -import torch - - -class ReproducibleL2Norm: - """Accumulate on a single device; reduce bins before calling ``finish``.""" - - def __init__(self, device: torch.device | None = None) -> None: - self.device = ( - device if device is not None else torch.device("cuda", torch.cuda.current_device()) - ) - - def cast(self, value: torch.Tensor, dtype: str) -> torch.Tensor: - """Convert to the named Torch dtype on the tensor's current device.""" - return value.to(getattr(torch, dtype)) - - def zeros(self, count: int = 23) -> torch.Tensor: - """Allocate zeroed integer bins on the accumulator device.""" - return torch.zeros(count, dtype=torch.int64, device=self.device) - - def tensor( - self, value: list[int | float] | int | float, dtype: str = "float32" - ) -> torch.Tensor: - """Create a typed constant on the accumulator device.""" - return torch.tensor(value, dtype=getattr(torch, dtype), device=self.device) - - def view(self, value: torch.Tensor, dtype: str) -> torch.Tensor: - """Interpret bits as the named Torch dtype without value conversion.""" - return value.view(getattr(torch, dtype)) - - def add(self, bins: torch.Tensor, indices: torch.Tensor, values: torch.Tensor) -> torch.Tensor: - """Add integer contributions to the selected bins in place.""" - return bins.scatter_add_(0, indices, values) - - def accumulate( - self, bins: torch.Tensor, gradient: torch.Tensor, chunk_size: int = 1048576 - ) -> torch.Tensor: - """Add device-computed FP32 gradient squares to exact integer bins.""" - if gradient.layout != torch.strided or gradient.dtype != torch.float32: - raise TypeError("Reproducible clipping requires FP32 gradients") - if chunk_size <= 0: - raise ValueError("chunk_size must be positive") - flat = gradient.reshape([-1]) - if flat.shape[0] > 2**40: - raise OverflowError("Reproducible norm supports at most 2**40 global elements") - bins = self.add(bins, self.tensor([22], "int64"), self.tensor([flat.shape[0]], "int64")) - for start in range(0, flat.shape[0], chunk_size): - value = flat[start : start + chunk_size] - bits = self.cast(self.view(value * value, "int32"), "int64") - exponent = (bits >> 23) & self.tensor(255, "int64") - fraction = bits & self.tensor(0x7FFFFF, "int64") - finite = exponent != 255 - mantissa = fraction | torch.where( - exponent > 0, torch.full_like(exponent, 0x800000), torch.zeros_like(exponent) - ) - shift = torch.maximum(exponent - 1, torch.zeros_like(exponent)) - limb = shift // 16 - shifted = (mantissa << (shift % 16)) * self.cast(finite, "int64") - for offset in range(3): - bins = self.add( - bins, limb + offset, (shifted >> (16 * offset)) & self.tensor(65535, "int64") - ) - flags = torch.stack( - [ - self.cast((exponent == 255) & (fraction == 0), "int64").sum(), - self.cast((exponent == 255) & (fraction != 0), "int64").sum(), - ] - ) - bins = self.add(bins, self.tensor([20, 21], "int64"), flags) - return bins - - def finish(self, bins: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - """Return norm and square sum after one FP32 rounding of reduced bins.""" - if int(bins[22].item()) > 2**40: - raise OverflowError("Reproducible norm supports at most 2**40 global elements") - carry = self.zeros(1)[0] - digits = [] - for i in range(20): - value = bins[i] + carry - digits.append(value & self.tensor(65535, "int64")) - carry = value >> 16 - digits = torch.stack(digits) - index = self.tensor(list(range(20)), "int64") - top = torch.where(digits != 0, index, torch.full_like(index, -1)).max() - - def get(i): - return torch.where(index == i, digits, torch.zeros_like(digits)).sum() - - word = get(top) - leading = self.zeros(1)[0] - for width in [8, 4, 2, 1]: - take = word >= (1 << width) - word = torch.where(take, word >> width, word) - leading = leading + self.cast(take, "int64") * width - highest = top * 16 + leading - cut = torch.maximum(highest - 23, self.zeros(1)[0]) - limb, shift = cut // 16, cut % 16 - significand = ( - (get(limb) >> shift) | (get(limb + 1) << (16 - shift)) | (get(limb + 2) << (32 - shift)) - ) & self.tensor(0xFFFFFF, "int64") - round_position = torch.maximum(cut - 1, self.zeros(1)[0]) - round_limb, round_shift = round_position // 16, round_position % 16 - round_word = get(round_limb) - round_bit = ((round_word >> round_shift) & self.tensor(1, "int64")) * self.cast( - cut > 0, "int64" - ) - sticky = torch.where(index < round_limb, digits, torch.zeros_like(digits)).sum() != 0 - sticky = sticky | ((round_word & ((self.tensor(1, "int64") << round_shift) - 1)) != 0) - significand = significand + round_bit * self.cast( - sticky | ((significand & self.tensor(1, "int64")) != 0), "int64" - ) - exponent = highest - 22 + self.cast(significand == 0x1000000, "int64") - raw = (exponent << 23) | (significand & self.tensor(0x7FFFFF, "int64")) - raw = torch.where(exponent >= 255, torch.full_like(raw, 0x7F800000), raw) - raw = torch.where(highest < 23, get(0) | (get(1) << 16), raw) - raw = torch.where(top < 0, torch.zeros_like(raw), raw) - raw = torch.where(bins[20] != 0, torch.full_like(raw, 0x7F800000), raw) - raw = torch.where(bins[21] != 0, torch.full_like(raw, 0x7FC00000), raw) - square_sum = self.view(self.cast(raw.reshape([1]), "int32"), "float32") - return torch.sqrt(square_sum), square_sum diff --git a/megatron/core/pipeline_parallel/schedules.py b/megatron/core/pipeline_parallel/schedules.py index b15b9e4a640..f576bc05a58 100644 --- a/megatron/core/pipeline_parallel/schedules.py +++ b/megatron/core/pipeline_parallel/schedules.py @@ -177,7 +177,11 @@ 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.uses_dsa_reference and config.tensor_model_parallel_size <= 1: + if ( + config is not None + and config.dsa_accuracy_compatible + and config.tensor_model_parallel_size <= 1 + ): return # Handle dict format (multi-module pipelines) @@ -571,7 +575,9 @@ def backward_step(input_tensor, output_tensor, output_tensor_grad, config): # 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.uses_dsa_reference or _tp_size > 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]) @@ -646,7 +652,7 @@ def _unwrap_single_tensor_list(tensor): 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.uses_dsa_reference or _tp_size > 1 + not config.dsa_accuracy_compatible or _tp_size > 1 ): custom_backward(output_tensor_module, output_tensor_grad_module) else: diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index ade60d8546e..b613c581ed2 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -349,7 +349,7 @@ def forward(self, 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, "uses_dsa_reference", False) and _tp_size <= 1: + 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] @@ -728,7 +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, - use_accuracy_compatible: bool = False, + dsa_accuracy_compatible: bool = False, ) -> torch.Tensor: """Linear layer execution with asynchronous communication and gradient accumulation fusion in backprop. @@ -795,7 +795,7 @@ 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 use_accuracy_compatible and _tp_size <= 1 and not sequence_parallel and not allreduce_dgrad: + 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 @@ -846,7 +846,7 @@ def _expert_grads_need_own_dp_domain(config) -> bool: """ if config.expert_model_parallel_size > 1: return True - if not getattr(config, "uses_dsa_reference", False): + if not getattr(config, 'dsa_accuracy_compatible', False): return False etp = getattr(config, 'expert_tensor_parallel_size', None) if etp is None: @@ -1080,7 +1080,7 @@ def _forward_impl(self, input, weight, *args, **kwargs): input, weight, *args, - use_accuracy_compatible=getattr(self.config, "uses_dsa_reference", False), + dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), **kwargs, ) @@ -1130,7 +1130,7 @@ def forward( or self.disable_grad_reduce ): input_parallel = input_ - elif getattr(self.config, "uses_dsa_reference", False) and ( + elif getattr(self.config, "dsa_accuracy_compatible", False) and ( self.tp_group is None or self.tp_group.size() <= 1 ): input_parallel = input_ @@ -1180,7 +1180,7 @@ def forward( gather_output = runtime_gather_output if gather_output and ( - not getattr(self.config, "uses_dsa_reference", False) + 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. @@ -1415,7 +1415,7 @@ def _forward_impl(self, input, weight, *args, **kwargs): input, weight, *args, - use_accuracy_compatible=getattr(self.config, "uses_dsa_reference", False), + dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), **kwargs, ) @@ -1467,7 +1467,7 @@ 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, "uses_dsa_reference", False) and ( + elif getattr(self.config, "dsa_accuracy_compatible", False) and ( self.tp_group is None or self.tp_group.size() <= 1 ): output_ = output_parallel diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index f93a779e790..8e0c135d725 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -121,7 +121,6 @@ class _AccuracyCompatibleSoftmax(torch.autograd.Function): @staticmethod def forward(ctx, logits: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor: - """Compute softmax over valid entries and retain its probabilities.""" 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) @@ -129,7 +128,6 @@ def forward(ctx, logits: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor @staticmethod def backward(ctx, grad_output: torch.Tensor): - """Apply the explicit softmax gradient and clear masked entries.""" probabilities, valid_mask = ctx.saved_tensors grad_logits = probabilities * ( grad_output - (grad_output * probabilities).sum(dim=-1, keepdim=True) @@ -154,7 +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, "uses_dsa_reference", False)) + 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: @@ -1831,7 +1829,7 @@ def forward( # 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.uses_dsa_reference and _tp_size <= 1): + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): x = x.detach() qr = qr.detach() diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index dc542b25433..67fc3106187 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -1352,7 +1352,7 @@ def forward( # 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.uses_dsa_reference: + if self.config.dsa_accuracy_compatible: pad_small_expert = ( self.config.use_accuracy_compatible and not self.config.moe_grouped_gemm diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index 9e34af3baca..798eb4cabd6 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -576,7 +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.uses_dsa_reference: + if self.config.dsa_accuracy_compatible: orig_dtype = output.dtype output = (output.float() + shared_expert_output.float()).to(orig_dtype) else: @@ -664,7 +664,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.uses_dsa_reference and not self.shared_expert_overlap: + if self.config.dsa_accuracy_compatible and not self.shared_expert_overlap: self._accuracy_shared_input = hidden_states_shared shared_expert_output = None else: @@ -703,7 +703,7 @@ 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.uses_dsa_reference: + 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) diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index bf1b4d892fe..6da367bb659 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -404,7 +404,7 @@ def permute( drop_and_pad: bool = False, tokens_per_expert: Optional[torch.Tensor] = None, align_size: int = 0, - use_accuracy_compatible: bool = False, + dsa_accuracy_compatible: bool = False, ) -> Tuple[ torch.Tensor, Optional[torch.Tensor], @@ -527,7 +527,7 @@ def permute( # === BIT-EXACT permute backward (gated by MOE_DETERMINISTIC_UNPERMUTE) === if ( _use_accuracy_compatible() - and not 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] diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index d12de4ad63d..a1ca6cdbe42 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -104,7 +104,7 @@ def gating(self, input: torch.Tensor): router_dtype = torch.float32 elif self.config.moe_router_dtype == 'fp64': router_dtype = torch.float64 - if self.config.uses_dsa_reference: + 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: diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 06ff5f22dfa..1b26b87e4e4 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -304,7 +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, - use_accuracy_compatible=self.config.uses_dsa_reference, + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) ) @@ -652,7 +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, - use_accuracy_compatible=self.config.uses_dsa_reference, + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) return permutated_local_input_tokens, permuted_probs @@ -1379,7 +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), - use_accuracy_compatible=self.config.uses_dsa_reference, + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) if self.router_dtype == "fp64": permuted_probs = permuted_probs.to(torch.float64) @@ -1529,7 +1529,7 @@ def token_dispatch( """ if self.shared_experts is not None: self.shared_experts.wait_current_stream() - if self.config.uses_dsa_reference: + if self.config.dsa_accuracy_compatible: async_finish = False allocate_on_comm_stream = False dispatched_hidden_states = self._comm_manager.dispatch( @@ -1591,7 +1591,7 @@ 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.uses_dsa_reference: + 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) diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 6ac4c47780b..5a91a7da037 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -606,7 +606,7 @@ def get_mtp_layer_spec_for_backend( column_parallel_linear_impl: type = backend.column_parallel_linear() layer_norm_impl = ( WrappedTorchNorm - if config is not None and config.uses_dsa_reference + if config is not None and config.norm_accuracy_compatible else backend.layer_norm() ) mtp_layer_spec = ModuleSpec( @@ -1116,7 +1116,7 @@ def _get_embeddings( decoder_input = decoder_input.detach() _tp_size = 1 if self.tp_group is None else self.tp_group.size() - if not (self.config.uses_dsa_reference and _tp_size <= 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 ) @@ -1136,12 +1136,12 @@ def _concat_embeddings(self, hidden_states: torch.Tensor, decoder_input: torch.T """ _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.uses_dsa_reference and _tp_size <= 1): + 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) - if not (self.config.uses_dsa_reference and _tp_size <= 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 ) @@ -1155,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 ) - elif not (self.config.uses_dsa_reference and _tp_size <= 1): + 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 8ea52ecd174..a7f9db5e60a 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -48,7 +48,7 @@ def __new__( assert not config.persist_layer_norm, f"persist_layer_norm not supported by torch LayerNorm" assert ( - config.uses_dsa_reference or not config.sequence_parallel + config.norm_accuracy_compatible or not config.sequence_parallel ), "sequence parallel not supported by torch LayerNorm" assert ( @@ -69,7 +69,7 @@ def __new__( raise Exception("Only LayerNorm, RMSNorm and L2Norm are currently supported") factory_kwargs = {} - if config.normalization == "RMSNorm" and config.uses_dsa_reference: + 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: diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index 1747e9406f8..b3b977f8053 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -591,7 +591,7 @@ def forward( # 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.uses_dsa_reference and _tp_size <= 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 ) @@ -699,7 +699,7 @@ def forward( # 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.uses_dsa_reference and _tp_size <= 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 ) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index a86da58f444..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"]}} ) @@ -325,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.""" @@ -1231,11 +1249,6 @@ class TransformerConfig(ModelParallelConfig): insert these joins. This feature is particularly useful when using with full-iteration CUDA graphs""" - @property - def uses_dsa_reference(self) -> bool: - """Select DSA reference numerics from the shared mode and model architecture.""" - return self.use_accuracy_compatible and self.experimental_attention_variant == "dsa" - def __post_init__(self): """Python dataclass method that is used to modify attributes after initialization. See https://docs.python.org/3/library/dataclasses.html#post-init-processing for more diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index 9af8bc96f83..5a584e59f42 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -941,7 +941,7 @@ 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.uses_dsa_reference and self.config.tensor_model_parallel_size <= 1: + if self.config.dsa_accuracy_compatible and self.config.tensor_model_parallel_size <= 1: output = hidden_states else: output = make_viewless_tensor( diff --git a/tests/unit_tests/distributed/test_finalize_model_grads.py b/tests/unit_tests/distributed/test_finalize_model_grads.py index 5ff0856ccef..762e1fdd81d 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -34,9 +34,8 @@ def test_token_normalization_preserves_legacy_and_dsa_contracts( 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, use_accuracy_compatible=accuracy + num_layers=1, hidden_size=8, num_attention_heads=1, dsa_accuracy_compatible=dsa ) - config.experimental_attention_variant = "dsa" if dsa else None gradient = torch.tensor(8.0, device="cuda") events = [] 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 5938247a4bd..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,7 +65,7 @@ def _make_config(**overrides): defaults = dict( num_layers=4, normalization="RMSNorm", - uses_dsa_reference=False, + norm_accuracy_compatible=False, qk_layernorm=False, multi_latent_attention=False, qk_l2_norm=False, @@ -380,7 +380,7 @@ def test_accuracy_compatible_qk_rmsnorm(self): qk_l2_norm=False, qk_layernorm=True, normalization="RMSNorm", - uses_dsa_reference=True, + norm_accuracy_compatible=True, ) spec = self._call(cfg=cfg, backend=backend) diff --git a/tests/unit_tests/optimizer/test_reproducible_norm.py b/tests/unit_tests/optimizer/test_reproducible_norm.py deleted file mode 100644 index 47e93092251..00000000000 --- a/tests/unit_tests/optimizer/test_reproducible_norm.py +++ /dev/null @@ -1,144 +0,0 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - -from types import SimpleNamespace - -import pytest -import torch - -import megatron.core.optimizer as optimizer_module -from megatron.core.optimizer.clip_grads import clip_grad_by_total_norm_fp32 -from megatron.core.optimizer.optimizer import ChainedOptimizer -from megatron.core.optimizer.optimizer_config import OptimizerConfig -from megatron.core.optimizer.reproducible_norm import ReproducibleL2Norm - - -@pytest.mark.parametrize( - 'values,expected', - [ - ([0.0], 0.0), - ([3.0, 4.0], 25.0), - ([4096.0, 1.0], 16777216.0), - ([4096.0, 1.0, 1.0, 1.0], 16777220.0), - ([2.0**-70], 2.0**-140), - ([float('inf')], float('inf')), - ([3e38], float('inf')), - ], -) -def test_sum_squares_rounding(values, expected): - norm = ReproducibleL2Norm() - _, squared = norm.finish(norm.accumulate(norm.zeros(), norm.tensor(values))) - assert squared.item() == expected - - -def test_nan_and_invalid_input(): - norm = ReproducibleL2Norm() - actual, _ = norm.finish( - norm.accumulate(norm.zeros(), norm.tensor([float('inf'), float('nan')])) - ) - assert torch.isnan(actual).item() - with pytest.raises(TypeError, match='FP32'): - norm.accumulate(norm.zeros(), norm.tensor([1.0]).bfloat16()) - bins = norm.zeros() - bins[22] = 2**40 + 1 - with pytest.raises(OverflowError): - norm.finish(bins) - - -def test_layout_and_chunk_invariance(): - norm = ReproducibleL2Norm() - gradient = torch.arange(1, 12289, device='cuda', dtype=torch.float32).reshape(96, 128) / 16384 - whole = norm.accumulate(norm.zeros(), gradient) - transposed = norm.accumulate(norm.zeros(), gradient.T.contiguous(), chunk_size=123) - split = norm.accumulate(norm.zeros(), gradient.flatten()[:1234]) - split += norm.accumulate(norm.zeros(), gradient.flatten()[1234:]) - assert torch.equal(whole, transposed) - assert torch.equal(whole, split) - - -def test_chained_owner_groups_and_clipping(monkeypatch): - rank = torch.distributed.get_rank() - world = torch.distributed.get_world_size() - singleton = None - for member in range(world): - group = torch.distributed.new_group([member]) - if member == rank: - singleton = group - config = OptimizerConfig(use_accuracy_compatible=True, clip_grad=1.0) - # Dense 3 and 4 have distinct owners; the expert 12 is replicated between - # singleton stats groups. Finishing each child separately loses this contract. - dense = [torch.tensor([3.0 if rank == 0 else 4.0], device='cuda')] if rank < 2 else [] - if world == 1: - dense = [torch.tensor([3.0, 4.0], device='cuda')] - expert = [torch.tensor([12.0], device='cuda')] - children = [ - SimpleNamespace( - config=config, - get_grads_for_grad_norm=lambda _group=None: dense, - get_grad_stats_parallel_group=lambda: torch.distributed.group.WORLD, - ), - SimpleNamespace( - config=config, - get_grads_for_grad_norm=lambda _group=None: expert, - get_grad_stats_parallel_group=lambda: singleton, - ), - ] - actual = ChainedOptimizer(children).get_grad_norm() - assert actual.item() == 13.0 - parameter = torch.nn.Parameter(torch.zeros(2, device='cuda')) - parameter.grad = torch.tensor([3.0, 4.0], device='cuda') - monkeypatch.setattr('megatron.core.optimizer.clip_grads.multi_tensor_scale_tensor_impl', None) - clip_grad_by_total_norm_fp32([parameter], 1.0, actual) - expected = torch.tensor([3.0, 4.0], device='cuda') * (1.0 / (actual + 1e-6)) - assert torch.equal(parameter.grad, expected) - - -@pytest.mark.parametrize('enabled,clip', [(False, 1.0), (True, 0.0)]) -def test_stock_norm_when_disabled_or_not_clipping(monkeypatch, enabled, clip): - config = OptimizerConfig(use_accuracy_compatible=enabled, clip_grad=clip) - optimizer = ChainedOptimizer([SimpleNamespace(config=config, get_grad_norm=lambda: 7.0)]) - - def unexpected(*args, **kwargs): - pytest.fail('Reproducible norm must not run without explicit enabled clipping') - - monkeypatch.setattr(optimizer, '_get_reproducible_grad_norm', unexpected) - assert optimizer.get_grad_norm() == 7.0 - - -@pytest.fixture(scope='module', autouse=True) -def distributed_norm_device(): - import os - - torch.cuda.set_device(int(os.environ.get('LOCAL_RANK', '0'))) - owns_group = not torch.distributed.is_initialized() - if owns_group: - torch.distributed.init_process_group(backend='nccl') - yield - if owns_group: - torch.distributed.destroy_process_group() - - -@pytest.fixture(scope='session') -def ensure_test_data(): - """The norm tests are self-contained and do not consume external datasets.""" - - -@pytest.mark.parametrize("enabled", [False, True]) -def test_optimizer_mode_selects_native_adam_without_secondary_flags(monkeypatch, enabled): - monkeypatch.setenv("USE_ACCURACY_COMPATIBLE", str(int(not enabled))) - config = OptimizerConfig(use_accuracy_compatible=enabled, lr=1e-3) - parameter = torch.nn.Parameter(torch.tensor([1.0], device="cuda")) - optimizer, _ = optimizer_module._get_megatron_optimizer_based_on_param_groups( - config, [], [{"params": [parameter]}], skip_megatron_wrapping=True - ) - if enabled: - assert type(optimizer) is torch.optim.AdamW - assert optimizer.defaults["fused"] is True - assert optimizer.defaults["foreach"] is False - else: - expected = ( - torch.optim.AdamW if optimizer_module.USING_PYTORCH_OPTIMIZER else optimizer_module.Adam - ) - assert type(optimizer) is expected - parameter.grad = torch.ones_like(parameter) - optimizer.step() - assert parameter.item() < 1.0 diff --git a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py index d53cd29ae36..979fff5282f 100644 --- a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py +++ b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py @@ -133,7 +133,7 @@ def test_embedding_configuration_controls_gradient_destination(self): weight.main_grad = torch.zeros_like(weight, dtype=torch.float32) instances.append( SimpleNamespace( - config=SimpleNamespace(uses_dsa_reference=enabled), + config=SimpleNamespace(dsa_accuracy_compatible=enabled), deterministic_mode=True, tp_group=_FakeGroup(1), weight=weight, @@ -141,7 +141,7 @@ def test_embedding_configuration_controls_gradient_destination(self): ) ) for instance in instances: - enabled = instance.config.uses_dsa_reference + enabled = instance.config.dsa_accuracy_compatible with patch.dict( os.environ, { @@ -261,7 +261,7 @@ def test_tp1_forward_dgrad_wgrad_bias_matches_f_linear(self): 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), use_accuracy_compatible=True + 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) @@ -294,12 +294,12 @@ def test_tp2_or_allreduce_skips_tp1_matmul_path(self): 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), use_accuracy_compatible=True + 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), use_accuracy_compatible=True + x, w, None, False, True, False, None, 0, _FakeGroup(1), dsa_accuracy_compatible=True ) self.assertIsNotNone(_SentinelApply.last) @@ -317,7 +317,7 @@ def test_deallocate_uac_tp1_preserves_storage(self): t = torch.arange(4.0, device="cuda", requires_grad=True) data_before = t.data.clone() deallocate_output_tensor( - t, True, SimpleNamespace(uses_dsa_reference=True, tensor_model_parallel_size=1) + 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,)) @@ -345,7 +345,7 @@ def test_backward_step_uac_tp1_uses_autograd_not_custom(self): grad_scale_func=None, deallocate_pipeline_outputs=True, tensor_model_parallel_size=1, - uses_dsa_reference=True, + dsa_accuracy_compatible=True, ) gin = backward_step(x, y, go, cfg) torch.testing.assert_close(gin, go * 3, atol=0, rtol=0) @@ -360,7 +360,7 @@ def test_backward_step_off_or_tp2_uses_custom_backward(self): grad_scale_func=None, deallocate_pipeline_outputs=True, tensor_model_parallel_size=1, - uses_dsa_reference=False, + dsa_accuracy_compatible=False, ) _UAC["on"] = False _CUSTOM_BWD["calls"] = [] @@ -371,7 +371,7 @@ def test_backward_step_off_or_tp2_uses_custom_backward(self): 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.uses_dsa_reference = True + cfg.dsa_accuracy_compatible = True cfg.tensor_model_parallel_size = 2 _UAC["on"] = True _CUSTOM_BWD["calls"] = [] diff --git a/tests/unit_tests/tensor_parallel/test_layers.py b/tests/unit_tests/tensor_parallel/test_layers.py index 57f5cea8ea9..be5795b2642 100644 --- a/tests/unit_tests/tensor_parallel/test_layers.py +++ b/tests/unit_tests/tensor_parallel/test_layers.py @@ -18,10 +18,10 @@ def test_expert_grads_need_own_dp_domain_etp_lt_tp(): expert_model_parallel_size=1, tensor_model_parallel_size=2, expert_tensor_parallel_size=1, - uses_dsa_reference=True, + dsa_accuracy_compatible=True, ) assert _expert_grads_need_own_dp_domain(frozen) is True - frozen.uses_dsa_reference = False + 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 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 db4c9450413..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 @@ -117,7 +117,7 @@ def capture(*args, **kwargs): key_positions=None, ) _run_sparse_attention(config=SimpleNamespace(), **common) - _run_sparse_attention(config=SimpleNamespace(uses_dsa_reference=True), **common) + _run_sparse_attention(config=SimpleNamespace(dsa_accuracy_compatible=True), **common) assert calls == [False, True] diff --git a/tests/unit_tests/transformer/moe/test_accuracy_migration.py b/tests/unit_tests/transformer/moe/test_accuracy_migration.py index cdd95668f86..620f71f44ad 100644 --- a/tests/unit_tests/transformer/moe/test_accuracy_migration.py +++ b/tests/unit_tests/transformer/moe/test_accuracy_migration.py @@ -87,7 +87,7 @@ def combine_postprocess(self, output): class _Flex: @property def config(self): - return SimpleNamespace(uses_dsa_reference=_UAC["on"]) + return SimpleNamespace(dsa_accuracy_compatible=_UAC["on"]) def __init__(self): self.shared_experts = None @@ -101,7 +101,7 @@ class _MoE: @property def config(self): return SimpleNamespace( - uses_dsa_reference=_UAC["on"], + dsa_accuracy_compatible=_UAC["on"], sequence_parallel=True, moe_shared_expert_overlap=False, moe_latent_size=0, diff --git a/tests/unit_tests/transformer/moe/test_routers.py b/tests/unit_tests/transformer/moe/test_routers.py index 1d18ad56d57..9f33079c9a9 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -67,8 +67,7 @@ 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.use_accuracy_compatible = True - self.router.config.experimental_attention_variant = "dsa" + self.router.config.router_accuracy_compatible = True logits = self.router.gating(hidden_states) expected = torch.mm( @@ -94,7 +93,7 @@ def fake_router_gating_linear(inp, weight, bias, router_dtype): ) hidden_states = torch.randn((3, 1, self.router.config.hidden_size), dtype=torch.bfloat16) - assert self.router.config.uses_dsa_reference is False + assert self.router.config.router_accuracy_compatible is False assert self.router.gating(hidden_states) is expected assert called diff --git a/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py b/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py index f348bb4a13c..ea0decdba51 100644 --- a/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py +++ b/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py @@ -38,9 +38,9 @@ def test_expert_gemm_rows_preserve_storage_contract(self, enabled, grouped, firs bias_activation_fusion=False, params_dtype=torch.bfloat16, use_accuracy_compatible=enabled, + dsa_accuracy_compatible=True, moe_grouped_gemm=grouped, ) - config.experimental_attention_variant = "dsa" experts = SequentialMLP( 2, config, diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index bbbec3f8e43..6f4c530b0c3 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -92,10 +92,9 @@ def test_accuracy_compatible_norms_override_te_mtp_norms(self): hidden_size=64, num_attention_heads=8, normalization="RMSNorm", - use_accuracy_compatible=True, + norm_accuracy_compatible=True, use_cpu_initialization=True, ) - config.experimental_attention_variant = "dsa" 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 diff --git a/tests/unit_tests/transformer/test_torch_norm.py b/tests/unit_tests/transformer/test_torch_norm.py index 140826b78a4..34d37ae4662 100644 --- a/tests/unit_tests/transformer/test_torch_norm.py +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -15,13 +15,11 @@ def _config(**overrides): "normalization": "RMSNorm", } values.update(overrides) - config = TransformerConfig(**values) - config.experimental_attention_variant = "dsa" - return config + return TransformerConfig(**values) def test_rmsnorm_uses_native_torch_implementation(): - config = _config(use_accuracy_compatible=True, params_dtype=torch.bfloat16) + 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) @@ -36,17 +34,7 @@ def test_sequence_parallel_is_opt_in_for_native_norm(): eps=1e-5, ) config = _config( - sequence_parallel=True, tensor_model_parallel_size=2, use_accuracy_compatible=True + 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()) - - -@pytest.mark.parametrize("enabled", [False, True]) -@pytest.mark.parametrize("variant", [None, "dsa"]) -def test_reference_mode_follows_one_switch_and_architecture(enabled, variant): - config = _config(use_accuracy_compatible=enabled) - config.experimental_attention_variant = variant - assert config.uses_dsa_reference is (enabled and variant == "dsa") - with pytest.raises(AttributeError): - config.uses_dsa_reference = True