From 941b153103cb1073d2a693e3adc0148b25e4a85a Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Mon, 14 Sep 2026 19:22:35 -0700 Subject: [PATCH 1/2] Add opt-in reproducible gradient clipping norm Preserve existing gradient ownership groups while making the FP32 norm independent of layout and partition order. Signed-off-by: Zhan Rongrui --- megatron/core/optimizer/clip_grads.py | 45 +++++- megatron/core/optimizer/optimizer.py | 36 ++++- megatron/core/optimizer/optimizer_config.py | 3 + megatron/core/optimizer/reproducible_norm.py | 135 ++++++++++++++++++ .../optimizer/test_reproducible_norm.py | 121 ++++++++++++++++ 5 files changed, 330 insertions(+), 10 deletions(-) create mode 100644 megatron/core/optimizer/reproducible_norm.py create mode 100644 tests/unit_tests/optimizer/test_reproducible_norm.py diff --git a/megatron/core/optimizer/clip_grads.py b/megatron/core/optimizer/clip_grads.py index 3c5491d39a1..7095beedcb4 100644 --- a/megatron/core/optimizer/clip_grads.py +++ b/megatron/core/optimizer/clip_grads.py @@ -50,12 +50,37 @@ 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, + reproducible_grad_norm: bool = False, ) -> float: """Calculate the p-norm of gradients in FP32 precision. @@ -80,6 +105,12 @@ def get_grad_norm_fp32( if isinstance(grads_for_norm, torch.Tensor): grads_for_norm = [grads_for_norm] + if reproducible_grad_norm: + 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) @@ -184,12 +215,14 @@ 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) - 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 - ) + 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() + ) 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 4a74328d0d9..3b516fb95ea 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -48,9 +48,15 @@ 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 +from .clip_grads import ( + clip_grad_by_total_norm_fp32, + count_zeros_fp32, + get_grad_norm_fp32, + get_reproducible_grad_norm_bins, +) from .grad_scaler import MegatronGradScaler from .optimizer_config import OptimizerConfig +from .reproducible_norm import ReproducibleL2Norm logger = getLogger(__name__) @@ -296,7 +302,9 @@ 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() + grads_for_norm, + grad_stats_parallel_group=self.get_grad_stats_parallel_group(), + reproducible_grad_norm=self.config.reproducible_grad_norm and self.config.clip_grad > 0, ) return total_norm @@ -308,7 +316,10 @@ 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() + grouped_grads, + grad_stats_parallel_group=self.get_grad_stats_parallel_group(), + reproducible_grad_norm=self.config.reproducible_grad_norm + and self.config.clip_grad > 0, ) self.grad_norms_by_group[grad_norm_group] = group_grad_norm return self.grad_norms_by_group @@ -326,7 +337,9 @@ 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() + grads_for_norm, + grad_stats_parallel_group=self.get_grad_stats_parallel_group(), + reproducible_grad_norm=self.config.reproducible_grad_norm and self.config.clip_grad > 0, ) if clip_grad > 0.0 and params: @@ -1566,8 +1579,21 @@ 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.reproducible_grad_norm 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(): @@ -1631,6 +1657,8 @@ 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.reproducible_grad_norm 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 a11f5c8cc3b..c7c846e93ad 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -371,6 +371,9 @@ class OptimizerConfig: ################ # Miscellaneous ################ + reproducible_grad_norm: bool = False + """Use a partition-independent FP32 gradient norm when clipping is enabled.""" + 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 new file mode 100644 index 00000000000..3380c3dbf84 --- /dev/null +++ b/megatron/core/optimizer/reproducible_norm.py @@ -0,0 +1,135 @@ +# 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/tests/unit_tests/optimizer/test_reproducible_norm.py b/tests/unit_tests/optimizer/test_reproducible_norm.py new file mode 100644 index 00000000000..14c3fefcf53 --- /dev/null +++ b/tests/unit_tests/optimizer/test_reproducible_norm.py @@ -0,0 +1,121 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from types import SimpleNamespace + +import pytest +import torch + +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(reproducible_grad_norm=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(reproducible_grad_norm=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.""" From 6d61b46159ea73c566763479eac7b3ca1d7b529e Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Tue, 15 Sep 2026 14:18:31 +0800 Subject: [PATCH 2/2] Consolidate reference numerics under the accuracy mode Signed-off-by: Zhan Rongrui --- .../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 | 4 +-- megatron/core/optimizer/optimizer.py | 12 +++++---- megatron/core/optimizer/optimizer_config.py | 7 ++--- 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 | 27 +++++++++++++++++-- .../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 ++++++++++--- 31 files changed, 128 insertions(+), 121 deletions(-) diff --git a/megatron/core/distributed/finalize_model_grads.py b/megatron/core/distributed/finalize_model_grads.py index d5b1e77066f..530824e1464 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.dsa_accuracy_compatible and num_tokens is not None + _use_accuracy_compatible() and not config.uses_dsa_reference 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.dsa_accuracy_compatible: + if loss_normalized_in_graph or config.uses_dsa_reference: 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 615e2c45196..955346196ef 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() and not self.config.uses_dsa_reference: 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( + print( # pylint: disable=bad-builtin 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 4b01ba561b6..7efa02576a2 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.norm_accuracy_compatible: + if rms_norm and config.uses_dsa_reference: 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 c7406191dc2..9d7e35e27e5 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 USING_PYTORCH_OPTIMIZER: + 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: 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,20 +587,7 @@ 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}) - # [对齐修复] 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) + 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 7095beedcb4..6345991e057 100644 --- a/megatron/core/optimizer/clip_grads.py +++ b/megatron/core/optimizer/clip_grads.py @@ -80,7 +80,7 @@ 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, - reproducible_grad_norm: bool = False, + use_accuracy_compatible: bool = False, ) -> float: """Calculate the p-norm of gradients in FP32 precision. @@ -105,7 +105,7 @@ def get_grad_norm_fp32( if isinstance(grads_for_norm, torch.Tensor): grads_for_norm = [grads_for_norm] - if reproducible_grad_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) diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index 3b516fb95ea..e00c17b200f 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -304,7 +304,8 @@ def get_grad_norm(self): total_norm = get_grad_norm_fp32( grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - reproducible_grad_norm=self.config.reproducible_grad_norm and self.config.clip_grad > 0, + use_accuracy_compatible=self.config.use_accuracy_compatible + and self.config.clip_grad > 0, ) return total_norm @@ -318,7 +319,7 @@ def _compute_grad_norms_by_group(self) -> Dict[str, float]: group_grad_norm = get_grad_norm_fp32( grouped_grads, grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - reproducible_grad_norm=self.config.reproducible_grad_norm + use_accuracy_compatible=self.config.use_accuracy_compatible and self.config.clip_grad > 0, ) self.grad_norms_by_group[grad_norm_group] = group_grad_norm @@ -339,7 +340,8 @@ def clip_grad_norm(self, clip_grad: float) -> float: grad_norm = get_grad_norm_fp32( grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - reproducible_grad_norm=self.config.reproducible_grad_norm and self.config.clip_grad > 0, + use_accuracy_compatible=self.config.use_accuracy_compatible + and self.config.clip_grad > 0, ) if clip_grad > 0.0 and params: @@ -1592,7 +1594,7 @@ def _get_reproducible_grad_norm(self, grad_norm_group=None): @torch.no_grad() def get_grad_norm(self): - if self.config.reproducible_grad_norm and self.config.clip_grad > 0: + 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() @@ -1657,7 +1659,7 @@ 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.reproducible_grad_norm and self.config.clip_grad > 0: + 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 = [] diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index c7c846e93ad..ba3c8a6aa10 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. @@ -371,8 +368,8 @@ class OptimizerConfig: ################ # Miscellaneous ################ - reproducible_grad_norm: bool = False - """Use a partition-independent FP32 gradient norm when clipping is enabled.""" + 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/pipeline_parallel/schedules.py b/megatron/core/pipeline_parallel/schedules.py index f576bc05a58..b15b9e4a640 100644 --- a/megatron/core/pipeline_parallel/schedules.py +++ b/megatron/core/pipeline_parallel/schedules.py @@ -177,11 +177,7 @@ 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 - ): + if config is not None and config.uses_dsa_reference and config.tensor_model_parallel_size <= 1: return # Handle dict format (multi-module pipelines) @@ -575,9 +571,7 @@ 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.dsa_accuracy_compatible or _tp_size > 1 - ): + if config.deallocate_pipeline_outputs and (not config.uses_dsa_reference 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]) @@ -652,7 +646,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.dsa_accuracy_compatible or _tp_size > 1 + not config.uses_dsa_reference 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 b613c581ed2..ade60d8546e 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, "dsa_accuracy_compatible", False) and _tp_size <= 1: + if getattr(self.config, "uses_dsa_reference", 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, - dsa_accuracy_compatible: bool = False, + use_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 dsa_accuracy_compatible and _tp_size <= 1 and not sequence_parallel and not allreduce_dgrad: + if use_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, 'dsa_accuracy_compatible', False): + if not getattr(config, "uses_dsa_reference", 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, - dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), + use_accuracy_compatible=getattr(self.config, "uses_dsa_reference", False), **kwargs, ) @@ -1130,7 +1130,7 @@ def forward( or self.disable_grad_reduce ): input_parallel = input_ - elif getattr(self.config, "dsa_accuracy_compatible", False) and ( + elif getattr(self.config, "uses_dsa_reference", 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, "dsa_accuracy_compatible", False) + not getattr(self.config, "uses_dsa_reference", 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, - dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), + use_accuracy_compatible=getattr(self.config, "uses_dsa_reference", 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, "dsa_accuracy_compatible", False) and ( + elif getattr(self.config, "uses_dsa_reference", 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 8e0c135d725..f93a779e790 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -121,6 +121,7 @@ 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) @@ -128,6 +129,7 @@ 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) @@ -152,7 +154,7 @@ def _run_sparse_attention( topk_length: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Run sparse attention for absorbed and non-absorbed MLA paths.""" - accuracy_compatible = bool(getattr(config, "dsa_accuracy_compatible", False)) + accuracy_compatible = bool(getattr(config, "uses_dsa_reference", False)) if absorbed_mla: latent_v_channels = int(getattr(config, "kv_lora_rank", 0) or 0) if latent_v_channels <= 0: @@ -1829,7 +1831,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.dsa_accuracy_compatible and _tp_size <= 1): + if not (self.config.uses_dsa_reference 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 67fc3106187..dc542b25433 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.dsa_accuracy_compatible: + if self.config.uses_dsa_reference: 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 798eb4cabd6..9e34af3baca 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.dsa_accuracy_compatible: + if self.config.uses_dsa_reference: 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.dsa_accuracy_compatible and not self.shared_expert_overlap: + if self.config.uses_dsa_reference 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.dsa_accuracy_compatible: + if self.config.uses_dsa_reference: 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 6da367bb659..bf1b4d892fe 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, - dsa_accuracy_compatible: bool = False, + use_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 dsa_accuracy_compatible + and not 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] diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index a1ca6cdbe42..d12de4ad63d 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.router_accuracy_compatible: + if self.config.uses_dsa_reference: 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 1b26b87e4e4..06ff5f22dfa 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, - dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, + use_accuracy_compatible=self.config.uses_dsa_reference, ) ) @@ -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, - dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, + use_accuracy_compatible=self.config.uses_dsa_reference, ) 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), - dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, + use_accuracy_compatible=self.config.uses_dsa_reference, ) 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.dsa_accuracy_compatible: + if self.config.uses_dsa_reference: 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.dsa_accuracy_compatible: + if self.config.uses_dsa_reference: 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 5a91a7da037..6ac4c47780b 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.norm_accuracy_compatible + if config is not None and config.uses_dsa_reference 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.dsa_accuracy_compatible and _tp_size <= 1): + if not (self.config.uses_dsa_reference 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.dsa_accuracy_compatible and _tp_size <= 1): + if not (self.config.uses_dsa_reference 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.dsa_accuracy_compatible and _tp_size <= 1): + if not (self.config.uses_dsa_reference 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.dsa_accuracy_compatible and _tp_size <= 1): + elif not (self.config.uses_dsa_reference 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 a7f9db5e60a..8ea52ecd174 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.norm_accuracy_compatible or not config.sequence_parallel + config.uses_dsa_reference 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.norm_accuracy_compatible: + if config.normalization == "RMSNorm" and config.uses_dsa_reference: 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 b3b977f8053..1747e9406f8 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.dsa_accuracy_compatible and _tp_size <= 1): + if not (self.config.uses_dsa_reference 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.dsa_accuracy_compatible and _tp_size <= 1): + if not (self.config.uses_dsa_reference 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 34c02443365..a86da58f444 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"]}} ) @@ -335,14 +325,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.""" @@ -1249,6 +1231,11 @@ 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 5a584e59f42..9af8bc96f83 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.dsa_accuracy_compatible and self.config.tensor_model_parallel_size <= 1: + if self.config.uses_dsa_reference 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 762e1fdd81d..5ff0856ccef 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -34,8 +34,9 @@ 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, dsa_accuracy_compatible=dsa + num_layers=1, hidden_size=8, num_attention_heads=1, use_accuracy_compatible=accuracy ) + 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 13cc7419c45..5938247a4bd 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", - norm_accuracy_compatible=False, + uses_dsa_reference=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", - norm_accuracy_compatible=True, + uses_dsa_reference=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 index 14c3fefcf53..47e93092251 100644 --- a/tests/unit_tests/optimizer/test_reproducible_norm.py +++ b/tests/unit_tests/optimizer/test_reproducible_norm.py @@ -5,6 +5,7 @@ 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 @@ -62,7 +63,7 @@ def test_chained_owner_groups_and_clipping(monkeypatch): group = torch.distributed.new_group([member]) if member == rank: singleton = group - config = OptimizerConfig(reproducible_grad_norm=True, clip_grad=1.0) + 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 [] @@ -93,7 +94,7 @@ def test_chained_owner_groups_and_clipping(monkeypatch): @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(reproducible_grad_norm=enabled, clip_grad=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): @@ -119,3 +120,25 @@ def distributed_norm_device(): @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 979fff5282f..d53cd29ae36 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(dsa_accuracy_compatible=enabled), + config=SimpleNamespace(uses_dsa_reference=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.dsa_accuracy_compatible + enabled = instance.config.uses_dsa_reference 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), dsa_accuracy_compatible=True + x, w, b, False, False, False, None, 0, _FakeGroup(1), use_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), dsa_accuracy_compatible=True + x, w, None, False, False, False, None, 0, _FakeGroup(2), use_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 + x, w, None, False, True, False, None, 0, _FakeGroup(1), use_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(dsa_accuracy_compatible=True, tensor_model_parallel_size=1) + t, True, SimpleNamespace(uses_dsa_reference=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, - dsa_accuracy_compatible=True, + uses_dsa_reference=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, - dsa_accuracy_compatible=False, + uses_dsa_reference=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.dsa_accuracy_compatible = True + cfg.uses_dsa_reference = 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 be5795b2642..57f5cea8ea9 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, - dsa_accuracy_compatible=True, + uses_dsa_reference=True, ) assert _expert_grads_need_own_dp_domain(frozen) is True - frozen.dsa_accuracy_compatible = False + frozen.uses_dsa_reference = 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 466eefe0b73..db4c9450413 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(dsa_accuracy_compatible=True), **common) + _run_sparse_attention(config=SimpleNamespace(uses_dsa_reference=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 620f71f44ad..cdd95668f86 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(dsa_accuracy_compatible=_UAC["on"]) + return SimpleNamespace(uses_dsa_reference=_UAC["on"]) def __init__(self): self.shared_experts = None @@ -101,7 +101,7 @@ class _MoE: @property def config(self): return SimpleNamespace( - dsa_accuracy_compatible=_UAC["on"], + uses_dsa_reference=_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 9f33079c9a9..1d18ad56d57 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -67,7 +67,8 @@ 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 + self.router.config.use_accuracy_compatible = True + self.router.config.experimental_attention_variant = "dsa" logits = self.router.gating(hidden_states) expected = torch.mm( @@ -93,7 +94,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.router_accuracy_compatible is False + assert self.router.config.uses_dsa_reference 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 ea0decdba51..f348bb4a13c 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 6f4c530b0c3..bbbec3f8e43 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -92,9 +92,10 @@ def test_accuracy_compatible_norms_override_te_mtp_norms(self): hidden_size=64, num_attention_heads=8, normalization="RMSNorm", - norm_accuracy_compatible=True, + use_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 34d37ae4662..140826b78a4 100644 --- a/tests/unit_tests/transformer/test_torch_norm.py +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -15,11 +15,13 @@ def _config(**overrides): "normalization": "RMSNorm", } values.update(overrides) - return TransformerConfig(**values) + config = TransformerConfig(**values) + config.experimental_attention_variant = "dsa" + return config def test_rmsnorm_uses_native_torch_implementation(): - config = _config(norm_accuracy_compatible=True, params_dtype=torch.bfloat16) + config = _config(use_accuracy_compatible=True, params_dtype=torch.bfloat16) norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) assert isinstance(norm, torch.nn.RMSNorm) @@ -34,7 +36,17 @@ def test_sequence_parallel_is_opt_in_for_native_norm(): eps=1e-5, ) config = _config( - sequence_parallel=True, tensor_model_parallel_size=2, norm_accuracy_compatible=True + sequence_parallel=True, tensor_model_parallel_size=2, use_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