modelscope / modelscope/ms-swift

GRPO训练DeepSpeed中报AssertionError

Open
#7,378 3 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

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

Description

Describe the bug
What the bug is, and how to reproduce, better with screenshots(描述bug以及复现过程,最好有截图)
External Mode下进行GRPO训练,1卡推理,7卡训练,每次都在同一个位置(一个epoch结束的位置)报错:
[rank2]: Traceback (most recent call last):
[rank2]: File "/home/HwHiAiUser/work/train_llm_env1.py", line 127, in
[rank2]: launch_train(args)
[rank2]: File "/home/HwHiAiUser/work/train_llm_env1.py", line 117, in launch_train
[rank2]: swiftRLHF(task_id, argv).main()
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/llm/base.py", line 49, in main
[rank2]: result = self.run()
[rank2]: File "/home/HwHiAiUser/work/trainfactory/grpo_ppo_workflow_env1.py", line 256, in run
[rank2]: self.train(trainer)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/llm/train/sft.py", line 225, in train
[rank2]: trainer.train(trainer.args.resume_from_checkpoint)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/mixin.py", line 675, in train
[rank2]: res = super().train(*args, **kwargs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/transformers/trainer.py", line 2245, in train
[rank2]: return inner_training_loop(
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/transformers/trainer.py", line 2560, in _inner_training_loop
[rank2]: tr_loss_step = self.training_step(model, inputs, num_items_in_batch)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/grpo_trainer.py", line 1617, in training_step
[rank2]: return super().training_step(model, inputs, num_items_in_batch)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/transformers/trainer.py", line 3746, in training_step
[rank2]: inputs = self._prepare_inputs(inputs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/utils.py", line 161, in wrapper
[rank2]: return func(self, *args, **kwargs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/grpo_trainer.py", line 471, in _prepare_inputs
[rank2]: generation_batch = self._generate_and_score_completions(generation_batch)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/grpo_trainer.py", line 1021, in _generate_and_score_completions
[rank2]: inputs = self._generate_completions(inputs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/grpo_trainer.py", line 992, in _generate_completions
[rank2]: inputs, outputs = self._fast_infer(inputs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/grpo_trainer.py", line 941, in _fast_infer
[rank2]: self._move_model_to_vllm()
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/utils.py", line 161, in wrapper
[rank2]: return func(self, *args, **kwargs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/swift/trainers/rlhf_trainer/grpo_trainer.py", line 624, in _move_model_to_vllm
[rank2]: with gather_if_zero3(parameters), patch_lora_merge(self.model, parameter_group):
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/deepspeed/runtime/zero/partition_parameters.py", line 2230, in exit
[rank2]: self.params[0].partition(param_list=self.params, has_been_updated=False)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/deepspeed/runtime/zero/partition_parameters.py", line 1375, in partition
[rank2]: self._partition(param_list, has_been_updated=has_been_updated)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/deepspeed/runtime/zero/partition_parameters.py", line 1524, in _partition
[rank2]: self._partition_param(param, has_been_updated=has_been_updated)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/deepspeed/utils/nvtx.py", line 15, in wrapped_fn
[rank2]: ret_val = func(*args, **kwargs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/deepspeed/runtime/zero/partition_parameters.py", line 1557, in _partition_param
[rank2]: free_param(param)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/deepspeed/utils/nvtx.py", line 15, in wrapped_fn
[rank2]: ret_val = func(*args, **kwargs)
[rank2]: File "/home/HwHiAiUser/miniconda3/envs/env1/lib/python3.10/site-packages/deepspeed/runtime/zero/partition_parameters.py", line 281, in free_param
[rank2]: assert not param.ds_active_sub_modules, param.ds_summary()
[rank2]: AssertionError: {'id': 0, 'status': 'AVAILABLE', 'numel': 622329856, 'ds_numel': 622329856, 'shape': (151936, 4096), 'ds_shape': (151936, 4096), 'requires_grad': False, 'grad_shape': None, 'persist': False, 'active_sub_modules': {4}, 'ds_tensor.shape': torch.Size([88904266])}

Your hardware and system info
Write your system info like CUDA version/system/GPU/torch version here(在这里给出硬件信息和系统信息,如CUDA版本,系统,GPU型号和torch版本等)
Ascend 910B3, torch/torch_npu=2.7.1, vllm/vllm_ascend=0.11.0

Additional context
Add any other context about the problem here(在这里补充其他信息)
Qwen3-8b模型,训练超参数: num_generation=4, steps_per_generation=4, gradient_accumulation=8, per_device_train_bs=1, num_process=7

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 at swift/trainers/rlhf_trainer/grpo_trainer.py, especially _fast_infer and _move_model_to_vllm, then trace the gather_if_zero3 and patch_lora_merge context in utils.py. Reproduce the Qwen3-8B GRPO setup with one inference card and seven training processes on Ascend 910B3, and determine why DeepSpeed retains an active submodule at epoch boundaries. Done means the run completes without the reported partition assertion.

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
38/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.