From c80cb24b4beeac202c2f3bc2eb5bae94a4cbfd9a Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Wed, 16 Sep 2026 17:56:22 +0800 Subject: [PATCH] Revert "Unify accuracy-mode configuration and preserve explicit clipping (#12)" This reverts commit 45908be019536a39881458fc592a24f8492c920c. --- swift/megatron/arguments/megatron_args.py | 1 + swift/megatron/init.py | 10 ++--- swift/megatron/pipelines/train/sft.py | 3 ++ swift/megatron/trainers/base.py | 3 -- swift/megatron/trainers/trainer.py | 6 +-- swift/megatron/utils/utils.py | 16 +++---- tests/megatron/test_accuracy_bridge_tp1.py | 2 +- tests/megatron/test_accuracy_loss_and_norm.py | 10 ++--- tests/megatron/test_model_config.py | 13 ++---- tests/megatron/test_reproducible_clipping.py | 44 ------------------- 10 files changed, 30 insertions(+), 78 deletions(-) delete mode 100644 tests/megatron/test_reproducible_clipping.py diff --git a/swift/megatron/arguments/megatron_args.py b/swift/megatron/arguments/megatron_args.py index b9a7dd0df9..0ff44c3efa 100644 --- a/swift/megatron/arguments/megatron_args.py +++ b/swift/megatron/arguments/megatron_args.py @@ -549,6 +549,7 @@ class MegatronArguments(RLHFMegatronArgumentsMixin, MegatronTunerMixin): start_weight_decay: Optional[float] = None end_weight_decay: Optional[float] = None clip_grad: float = 1. + native_unfused_adamw: bool = False adam_beta1: float = 0.9 adam_beta2: float = 0.95 adam_eps: float = 1e-8 diff --git a/swift/megatron/init.py b/swift/megatron/init.py index 433faf36f9..4a5e173af2 100644 --- a/swift/megatron/init.py +++ b/swift/megatron/init.py @@ -183,7 +183,6 @@ def _patch_mcore_bridge_disable_te(): import mcore_bridge.model.register as mcb_register def _force_local_spec(orig): - def wrapper(*args, **kwargs): kwargs['use_transformer_engine'] = False return orig(*args, **kwargs) @@ -202,7 +201,7 @@ def wrapper(*args, **kwargs): origin_backend_spec_provider = _eav._get_backend_spec_provider def _local_backend_spec_provider(config): - if not config.uses_dsa_reference: + if not getattr(config, 'dsa_accuracy_compatible', False): return origin_backend_spec_provider(config) from megatron.core.models.backends import LocalSpecProvider return LocalSpecProvider() @@ -274,7 +273,8 @@ def scoped(method): @wraps(method) def call(self, *args, **kwargs): config = self.config - token = active.set(config.uses_dsa_reference and config.tensor_model_parallel_size <= 1) + token = active.set( + getattr(config, 'dsa_accuracy_compatible', False) and config.tensor_model_parallel_size <= 1) try: return method(self, *args, **kwargs) finally: @@ -327,7 +327,7 @@ def replace_spec_dsa(self, layer_spec): from megatron.core.transformer.torch_norm import WrappedTorchNorm dsa_spec = layer_spec.submodules.self_attention - if self.config.uses_dsa_reference: + if getattr(self.config, 'norm_accuracy_compatible', False): dsa_spec.submodules.q_layernorm = WrappedTorchNorm dsa_spec.submodules.kv_layernorm = WrappedTorchNorm indexer = getattr( @@ -335,7 +335,7 @@ def replace_spec_dsa(self, layer_spec): 'indexer', None, ) - if (self.config.uses_dsa_reference and indexer is not None + if (getattr(self.config, 'norm_accuracy_compatible', False) and indexer is not None and getattr(indexer, 'submodules', None) is not None): indexer.submodules.k_norm = WrappedTorchNorm diff --git a/swift/megatron/pipelines/train/sft.py b/swift/megatron/pipelines/train/sft.py index 9b007a2f14..381a02b24d 100644 --- a/swift/megatron/pipelines/train/sft.py +++ b/swift/megatron/pipelines/train/sft.py @@ -43,6 +43,9 @@ def __init__(self, args: Optional[Union[List[str], MegatronSftArguments]] = None self.train_msg = {} super(SwiftSft, self).__init__(args) args = self.args + from megatron.core.transformer.module import _use_accuracy_compatible + if _use_accuracy_compatible() and getattr(args, 'clip_grad', 0) > 0: + args.clip_grad = 0.0 if repatch is not None: megatron_args = asdict(self.args) if args.attention_backend != 'local': diff --git a/swift/megatron/trainers/base.py b/swift/megatron/trainers/base.py index 53705abc88..0b0d5da9ad 100644 --- a/swift/megatron/trainers/base.py +++ b/swift/megatron/trainers/base.py @@ -222,9 +222,6 @@ def get_optimizer_and_scheduler(self): else: config_cls = OptimizerConfig - if args.use_accuracy_compatible and not hasattr(config_cls, 'use_accuracy_compatible'): - raise ValueError('use_accuracy_compatible requires a Megatron-Core version that supports it') - kwargs = { f.name: getattr(args, f.name) for f in dataclasses.fields(config_cls) if hasattr(args, f.name) and f.name != 'loss_scale' diff --git a/swift/megatron/trainers/trainer.py b/swift/megatron/trainers/trainer.py index 46b8c8f2cb..603fafead1 100644 --- a/swift/megatron/trainers/trainer.py +++ b/swift/megatron/trainers/trainer.py @@ -179,9 +179,9 @@ def loss_func(self, import hashlib as _hashlib _final = (loss[0].detach().float() / loss[1].detach().float().clamp(min=1)).contiguous() print( - f'\nfinal_loss: rank={torch.distributed.get_rank()} ' - f'val={_final.item():.20f} ' - f'md5={_hashlib.md5(_final.cpu().numpy().tobytes()).hexdigest()}', + f"\nfinal_loss: rank={torch.distributed.get_rank()} " + f"val={_final.item():.20f} " + f"md5={_hashlib.md5(_final.cpu().numpy().tobytes()).hexdigest()}", flush=True) metrics = {'loss': reporting_loss} diff --git a/swift/megatron/utils/utils.py b/swift/megatron/utils/utils.py index b7e760ad2a..4774c99095 100644 --- a/swift/megatron/utils/utils.py +++ b/swift/megatron/utils/utils.py @@ -213,7 +213,7 @@ def get_padding_to(args): if args.tensor_model_parallel_size > 1 and args.sequence_parallel: padding_to = args.tensor_model_parallel_size # Match the DSA reference carrier without changing other TP+SP models. - if getattr(args, 'use_accuracy_compatible', False) and args.model_type == 'glm_moe_dsa': + if (getattr(args, 'megatron_extra_kwargs', None) or {}).get('dsa_accuracy_compatible', False): padding_to *= 2 if args.context_parallel_size > 1: padding_to = (padding_to or 1) * args.context_parallel_size @@ -292,7 +292,7 @@ def get_load_fixed_data_path(): def _batch_data_suffix(step, rank, seq_len): - return f'step{step}_rank{rank}_seq{seq_len}.npy' + return f"step{step}_rank{rank}_seq{seq_len}.npy" def dump_batch_data(batch, step, seq_len): @@ -307,10 +307,10 @@ def dump_batch_data(batch, step, seq_len): torch.cuda.synchronize() os.makedirs(dump_path, exist_ok=True) suffix = _batch_data_suffix(step, rank, seq_len) - np.save(os.path.join(dump_path, f'tokens_{suffix}'), tokens.detach().cpu().numpy()) - np.save(os.path.join(dump_path, f'labels_{suffix}'), labels.detach().cpu().numpy()) + np.save(os.path.join(dump_path, f"tokens_{suffix}"), tokens.detach().cpu().numpy()) + np.save(os.path.join(dump_path, f"labels_{suffix}"), labels.detach().cpu().numpy()) if rank == 0: - print(f'[DUMP_DATA_PATH] saved tokens_{suffix} and labels_{suffix}', flush=True) + print(f"[DUMP_DATA_PATH] saved tokens_{suffix} and labels_{suffix}", flush=True) def load_fixed_batch_data(batch, step, seq_len): @@ -320,11 +320,11 @@ def load_fixed_batch_data(batch, step, seq_len): rank = torch.distributed.get_rank() if torch.distributed.is_initialized() else 0 suffix = _batch_data_suffix(step, rank, seq_len) - tokens_file = os.path.join(load_path, f'tokens_{suffix}') - labels_file = os.path.join(load_path, f'labels_{suffix}') + tokens_file = os.path.join(load_path, f"tokens_{suffix}") + labels_file = os.path.join(load_path, f"labels_{suffix}") if not (os.path.exists(tokens_file) and os.path.exists(labels_file)): if rank == 0: - print(f'[LOAD_FIXED_DATA_PATH] file not found: {tokens_file}', flush=True) + print(f"[LOAD_FIXED_DATA_PATH] file not found: {tokens_file}", flush=True) return batch tokens_np = np.load(tokens_file) diff --git a/tests/megatron/test_accuracy_bridge_tp1.py b/tests/megatron/test_accuracy_bridge_tp1.py index 0aeed1edf0..9c8ed89768 100644 --- a/tests/megatron/test_accuracy_bridge_tp1.py +++ b/tests/megatron/test_accuracy_bridge_tp1.py @@ -45,7 +45,7 @@ def forward(self, inp, callback=None, module=module): def instance(self, cls, enabled, tp_size=1): instance = cls() - instance.config = SimpleNamespace(uses_dsa_reference=enabled, tensor_model_parallel_size=tp_size) + instance.config = SimpleNamespace(dsa_accuracy_compatible=enabled, tensor_model_parallel_size=tp_size) return instance def test_idempotent(self): diff --git a/tests/megatron/test_accuracy_loss_and_norm.py b/tests/megatron/test_accuracy_loss_and_norm.py index e02171a706..235f94ab32 100644 --- a/tests/megatron/test_accuracy_loss_and_norm.py +++ b/tests/megatron/test_accuracy_loss_and_norm.py @@ -81,8 +81,8 @@ def test_indexer_norm_preserves_disabled_provider(self): provider_norm = type('ProviderNorm', (), {}) module = types.ModuleType('megatron.core.transformer.torch_norm') module.WrappedTorchNorm = native_norm - for enabled in (False, True): - with self.subTest(accuracy=enabled): + for enabled, norm_accuracy in ((False, False), (True, False), (True, True)): + with self.subTest(accuracy=enabled, norm_accuracy=norm_accuracy): calls = [] replace_spec = production_function( 'swift/megatron/init.py', 'replace_spec_dsa', { @@ -96,12 +96,12 @@ def test_indexer_norm_preserves_disabled_provider(self): kv_layernorm=provider_norm, core_attention=types.SimpleNamespace(submodules=types.SimpleNamespace(indexer=indexer)))) spec = types.SimpleNamespace(submodules=types.SimpleNamespace(self_attention=attention)) - loader = types.SimpleNamespace(config=types.SimpleNamespace(uses_dsa_reference=enabled)) + loader = types.SimpleNamespace(config=types.SimpleNamespace(norm_accuracy_compatible=norm_accuracy)) with patch.dict(sys.modules, {module.__name__: module}): replace_spec(loader, spec) self.assertEqual(calls, ['provider']) - self.assertIs(indexer.submodules.k_norm, native_norm if enabled else provider_norm) - expected_qkv = native_norm if enabled else provider_norm + self.assertIs(indexer.submodules.k_norm, native_norm if norm_accuracy else provider_norm) + expected_qkv = native_norm if norm_accuracy else provider_norm self.assertIs(attention.submodules.q_layernorm, expected_qkv) self.assertIs(attention.submodules.kv_layernorm, expected_qkv) diff --git a/tests/megatron/test_model_config.py b/tests/megatron/test_model_config.py index c4341d1a93..bbf85f99d0 100644 --- a/tests/megatron/test_model_config.py +++ b/tests/megatron/test_model_config.py @@ -103,7 +103,7 @@ def test_get_mcore_model_config_does_not_enable_mtp_from_nested_checkpoint(monke def test_get_mcore_model_config_prefers_n_routed_experts(monkeypatch): _patch_model_config(monkeypatch) - hf_config = PretrainedConfig(model_type='glm_moe_dsa', num_experts=256, n_routed_experts=16) + hf_config = PretrainedConfig(model_type="glm_moe_dsa", num_experts=256, n_routed_experts=16) config = utils.get_mcore_model_config(_make_args(), hf_config) @@ -130,16 +130,11 @@ def test_get_padding_to_sequence_parallel_uses_tp_times_two(): fp4_format=None, fp4=None, attention_backend='unfused', - model_type='glm_moe_dsa', - model_info=SimpleNamespace(config=None), ) assert get_padding_to(args) == 2 - args.use_accuracy_compatible = True + args.megatron_extra_kwargs = {"dsa_accuracy_compatible": True} assert get_padding_to(args) == 4 - args.model_type = 'glm4_moe' - assert get_padding_to(args) == 2 - args.model_type = 'glm_moe_dsa' - args.use_accuracy_compatible = False + args.megatron_extra_kwargs = {"dsa_accuracy_compatible": False} assert get_padding_to(args) == 2 seq_len = 57 assert math.ceil(seq_len / 4) * 4 == 60 @@ -164,7 +159,7 @@ def test_dsa_backend_forced_to_local_spec_when_accuracy_compatible(monkeypatch): monkeypatch.setattr(init, '_use_accuracy_compatible_enabled', lambda: True) init._patch_mcore_bridge_disable_te() - provider = eav._get_backend_spec_provider(SimpleNamespace(uses_dsa_reference=True)) + provider = eav._get_backend_spec_provider(SimpleNamespace(dsa_accuracy_compatible=True)) assert isinstance(provider, LocalSpecProvider) assert hasattr(provider, 'linear') assert provider.linear() is not provider.column_parallel_linear() diff --git a/tests/megatron/test_reproducible_clipping.py b/tests/megatron/test_reproducible_clipping.py deleted file mode 100644 index 4afb5aeb53..0000000000 --- a/tests/megatron/test_reproducible_clipping.py +++ /dev/null @@ -1,44 +0,0 @@ -"""Explicit clipping survives accuracy-mode initialization.""" - -import pytest -from types import SimpleNamespace -from unittest.mock import patch - -from swift.megatron.pipelines.train import sft as sft_module -from swift.megatron.trainers import base as trainer_module -from swift.pipelines.base import SwiftPipeline - - -@pytest.mark.parametrize('accuracy', [False, True]) -@pytest.mark.parametrize('clip', [0.0, 0.25, 1.0]) -def test_constructor_preserves_explicit_clipping(accuracy, clip): - args = SimpleNamespace( - clip_grad=clip, - use_accuracy_compatible=accuracy, - template_meta=SimpleNamespace(template_cls=None), - model_meta=SimpleNamespace(is_multimodal=False), - mcore_model='existing-model', - output_dir='unused', - get_model_processor=lambda **kwargs: (None, None), - save_args=lambda output_dir: None) - - def pipeline_init(instance, values): - instance.args = values - - def prepare_template(instance): - instance.template = SimpleNamespace() - - with patch.object(SwiftPipeline, '__init__', pipeline_init), \ - patch.object(sft_module, 'repatch', None), \ - patch.object(sft_module.MegatronSft, '_prepare_template', prepare_template), \ - patch('megatron.core.transformer.module._use_accuracy_compatible', return_value=accuracy): - instance = sft_module.MegatronSft(args) - assert instance.args.clip_grad == clip - - -def test_unsupported_megatron_rejects_requested_norm(): - trainer = SimpleNamespace(args=SimpleNamespace(use_accuracy_compatible=True)) - with patch.object(trainer_module, 'mcore_016', False), \ - patch.object(trainer_module, 'OptimizerConfig', type('OldConfig', (), {})), \ - pytest.raises(ValueError, match='Megatron-Core version'): - trainer_module.BaseMegatronTrainer.get_optimizer_and_scheduler(trainer)