modelscope / modelscope/ms-swift

关于训练速度越来越快的疑问以及为何在接近结束时会疑似因为显存不够报错

Open
#9,681 2 comments 1 reaction 0 assignees View on GitHub

Nobody has claimed this yet.

question
Dominant language
Python
Stars
15.7k
Forks
1.7k
Avg merge
1d 16h
Merged PRs (30d)
136

Description

Checklist / 检查清单
  • I have searched existing issues, and this is a new question or discussion topic. / 我已经搜索过现有的 issues,确认这是一个新的问题与讨论。
Question Description / 问题描述

Train: 1%| | 5/508 [03:41<6:00:04, 42.95s/it]
Train: 1%| | 6/508 [04:24<5:58:12, 42.81s/it]
Train: 1%|▏ | 7/508 [05:04<5:50:09, 41.94s/it]
Train: 2%|▏ | 8/508 [05:45<5:46:41, 41.60s/it]
Train: 2%|▏ | 9/508 [06:25<5:41:20, 41.04s/it]

Train: 21%|██ | 106/508 [1:10:49<4:17:09, 38.38s/it]
Train: 21%|██ | 107/508 [1:11:24<4:09:31, 37.33s/it]
Train: 21%|██▏ | 108/508 [1:11:58<4:02:38, 36.40s/it]
Train: 21%|██▏ | 109/508 [1:12:33<4:00:26, 36.16s/it]
Train: 22%|██▏ | 110/508 [1:13:08<3:57:00, 35.73s/it]

Train: 40%|████ | 205/508 [2:08:32<3:00:04, 35.66s/it]
Train: 41%|████ | 206/508 [2:08:57<2:43:20, 32.45s/it]
Train: 41%|████ | 207/508 [2:09:31<2:44:17, 32.75s/it]
Train: 41%|████ | 208/508 [2:10:08<2:50:56, 34.19s/it]
Train: 41%|████ | 209/508 [2:10:38<2:43:08, 32.74s/it]

Train: 59%|█████▉ | 301/508 [2:53:34<1:42:06, 29.60s/it]
Train: 59%|█████▉ | 302/508 [2:53:49<1:26:18, 25.14s/it]
Train: 60%|█████▉ | 303/508 [2:54:14<1:25:42, 25.08s/it]
Train: 60%|█████▉ | 304/508 [2:54:45<1:31:47, 27.00s/it]
Train: 60%|██████ | 305/508 [2:55:16<1:34:54, 28.05s/it]

Train: 81%|████████ | 411/508 [3:37:04<39:40, 24.54s/it]
Train: 81%|████████ | 412/508 [3:37:21<35:53, 22.43s/it]
Train: 81%|████████▏ | 413/508 [3:37:50<38:32, 24.34s/it]
Train: 81%|████████▏ | 414/508 [3:38:10<35:49, 22.87s/it]
Train: 82%|████████▏ | 415/508 [3:38:35<36:35, 23.61s/it]

[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/cli/sft.py", line 20, in
[rank1]: sft_main()
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/pipelines/train/sft.py", line 340, in sft_main
[rank1]: return SwiftSft(args).main()
[rank1]: ^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/pipelines/base.py", line 52, in main
[rank1]: result = self.run()
[rank1]: ^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/ray_utils/base.py", line 168, in wrapper
[rank1]: return func(self, *args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/pipelines/train/sft.py", line 184, in run
[rank1]: return self.train(trainer)
[rank1]: ^^^^^^^^^^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/pipelines/train/sft.py", line 258, in train
[rank1]: trainer.train(resume_checkpoint)
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/trainers/mixin.py", line 967, in train
[rank1]: res = super().train(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/trainer.py", line 1433, in train
[rank1]: return inner_training_loop(
[rank1]: ^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/trainer.py", line 1515, in _inner_training_loop
[rank1]: self._run_epoch(
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/trainer.py", line 1743, in _run_epoch
[rank1]: tr_loss_step = self.training_step(model, inputs, num_items_in_batch)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/trainers/seq2seq_trainer.py", line 233, in training_step
[rank1]: return super().training_step(model, inputs, *args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/trainer.py", line 1915, in training_step
[rank1]: loss = self.compute_loss(model, inputs, num_items_in_batch=num_items_in_batch)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/trainers/seq2seq_trainer.py", line 141, in compute_loss
[rank1]: outputs = self.template.compute_sft_loss(model, inputs, num_items_in_batch=num_items_in_batch, trainer=self)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/template/base.py", line 774, in compute_sft_loss
[rank1]: outputs = model(**inputs)
[rank1]: ^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
[rank1]: return self._call_impl(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
[rank1]: return forward_call(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/deepspeed/utils/nvtx.py", line 20, in wrapped_fn
[rank1]: ret_val = func(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/deepspeed/runtime/engine.py", line 2358, in forward
[rank1]: loss = self.module(*inputs, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
[rank1]: return self._call_impl(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1885, in _call_impl
[rank1]: return inner()
[rank1]: ^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1833, in inner
[rank1]: result = forward_call(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/utils/generic.py", line 907, in wrapper
[rank1]: output = func(self, *args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 492, in forward
[rank1]: outputs: BaseModelOutputWithPast = self.model(
[rank1]: ^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
[rank1]: return self._call_impl(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
[rank1]: return forward_call(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/utils/generic.py", line 1036, in wrapper
[rank1]: output = func(self, *args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/utils/output_capturing.py", line 252, in wrapper
[rank1]: outputs = func(self, *args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 424, in forward
[rank1]: hidden_states = decoder_layer(
[rank1]: ^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/modeling_layers.py", line 92, in call
[rank1]: return self._gradient_checkpointing_func(partial(super().call, **kwargs), *args)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/data/xxx/project/rec/thirdparty/ms-swift/swift/trainers/mixin.py", line 899, in _new_checkpoint
[rank1]: return old_checkpoint(*args, use_reentrant=use_reentrant, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/_compile.py", line 54, in inner
[rank1]: return disable_fn(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py", line 1263, in _fn
[rank1]: return fn(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/utils/checkpoint.py", line 505, in checkpoint
[rank1]: return CheckpointFunction.apply(function, preserve, *args)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/autograd/function.py", line 596, in apply
[rank1]: return super().apply(*args, **kwargs) # type: ignore[misc]
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/utils/checkpoint.py", line 268, in forward
[rank1]: outputs = run_function(*args)
[rank1]: ^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
[rank1]: return self._call_impl(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
[rank1]: return forward_call(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 318, in forward
[rank1]: hidden_states, _ = self.self_attn(
[rank1]: ^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1779, in _wrapped_call_impl
[rank1]: return self._call_impl(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1790, in _call_impl
[rank1]: return forward_call(*args, **kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 277, in forward
[rank1]: attn_output, attn_weights = attention_interface(
[rank1]: ^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xxx/.conda/envs/recllm/lib/python3.12/site-packages/transformers/integrations/sdpa_attention.py", line 92, in sdpa_attention_forward
[rank1]: attn_output = torch.nn.functional.scaled_dot_product_attention(
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: RuntimeError: Expected mha_graph.execute(handle, variant_pack, workspace_ptr.get()).is_good() to be true, but got false. (Could this error message be improved? If so, please report an enhancement request to PyTorch.)
W0704 08:15:11.465000 488675 site-packages/torch/distributed/elastic/multiprocessing/api.py:1012] Sending process 488756 closing signal SIGTERM
E0704 08:15:13.594000 488675 site-packages/torch/distributed/elastic/multiprocessing/api.py:986] failed (exitcode: 1) local_rank: 1 (pid: 488757) of binary: /home/xxx/.conda/envs/recllm/bin/python3.12

在使用swift sft的过程中发现两个问题,一个是为什么训练的速度会越来越快,因为训练是两个epoch,之前没有遇到过第二个epoch比第一个epoch训练速度快这么多的情况,第二个问题似乎是由于显存不够引起的,但很奇怪为什么会到训练到92%的时候才出现,且该报错几次都出现在训练最后快结束时

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with the reported traceback in swift/pipelines/train/sft.py, swift/trainers/mixin.py, swift/trainers/seq2seq_trainer.py, and swift/template/base.py, then follow the call into PyTorch scaled dot-product attention. Reproduce the training failure described in the issue and determine whether the speed change and final runtime error are related. Done means documenting a confirmed cause and actionable resolution or a narrowly scoped project fix.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems, machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Quiet
Clarity
Needs clarification
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.