Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions swift/megatron/arguments/megatron_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
10 changes: 5 additions & 5 deletions swift/megatron/init.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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()
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -327,15 +327,15 @@ 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(
getattr(dsa_spec.submodules.core_attention, 'submodules', None),
'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

Expand Down
3 changes: 3 additions & 0 deletions swift/megatron/pipelines/train/sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -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':
Expand Down
3 changes: 0 additions & 3 deletions swift/megatron/trainers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'
Expand Down
6 changes: 3 additions & 3 deletions swift/megatron/trainers/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down
16 changes: 8 additions & 8 deletions swift/megatron/utils/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand All @@ -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):
Expand All @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion tests/megatron/test_accuracy_bridge_tp1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
10 changes: 5 additions & 5 deletions tests/megatron/test_accuracy_loss_and_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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', {
Expand All @@ -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)

Expand Down
13 changes: 4 additions & 9 deletions tests/megatron/test_model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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
Expand All @@ -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()
Expand Down
44 changes: 0 additions & 44 deletions tests/megatron/test_reproducible_clipping.py

This file was deleted.

Loading