modelscope / modelscope/ms-swift
qwen3.6-35b DPO开启序列并行 (--sequence_parallel_size 4)报错: torch.AcceleratorError: CUDA error: an illegal memory access was encountered
Nobody has claimed this yet.
- 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 bug report. / 我已经搜索过现有的 issues,确认这是一个新的 bug report。
Bug Description / Bug 描述
在训练DPO时开启序列并行 (--sequence_parallel_size 4)报错:
torch.AcceleratorError: CUDA error: an illegal memory access was encountered
使用qwen3.6-27b训练正常,qwen3.6-36b-a3b会报错。
完整报错如下:
[rank0]: Traceback (most recent call last):
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/cli/rlhf.py", line 7, in <module>
[rank0]: rlhf_main()
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/pipelines/train/rlhf.py", line 246, in rlhf_main
[rank0]: return SwiftRLHF(args).main()
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/pipelines/base.py", line 52, in main
[rank0]: result = self.run()
[rank0]: ^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/ray_utils/base.py", line 168, in wrapper
[rank0]: return func(self, *args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/pipelines/train/sft.py", line 197, in run
[rank0]: return self.train(trainer)
[rank0]: ^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/pipelines/train/sft.py", line 271, in train
[rank0]: trainer.train(resume_checkpoint)
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/trainers/mixin.py", line 914, in train
[rank0]: res = super().train(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/trainer.py", line 1433, in train
[rank0]: return inner_training_loop(
[rank0]: ^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/trainer.py", line 1515, in _inner_training_loop
[rank0]: self._run_epoch(
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/trainer.py", line 1743, in _run_epoch
[rank0]: tr_loss_step = self.training_step(model, inputs, num_items_in_batch)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/rlhf_trainers/dpo_trainer.py", line 462, in training_step
[rank0]: return super().training_step(model, inputs, *args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/trl/trainer/dpo_trainer.py", line 1449, in training_step
[rank0]: return super().training_step(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/trainer.py", line 1915, in training_step
[rank0]: loss = self.compute_loss(model, inputs, num_items_in_batch=num_items_in_batch)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/rlhf_trainers/dpo_trainer.py", line 448, in compute_loss
[rank0]: loss, metrics = self.get_batch_loss_metrics(model, inputs, train_eval='train')
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/rlhf_trainers/dpo_trainer.py", line 366, in get_batch_loss_metrics
[rank0]: model_output = self.concatenated_forward(model, batch)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/rlhf_trainers/dpo_trainer.py", line 119, in concatenated_forward
[rank0]: outputs = model(**batch, use_cache=False)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1776, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1787, in _call_impl
[rank0]: return forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/deepspeed/utils/nvtx.py", line 20, in wrapped_fn
[rank0]: ret_val = func(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/deepspeed/runtime/engine.py", line 2358, in forward
[rank0]: loss = self.module(*inputs, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1776, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1882, in _call_impl
[rank0]: return inner()
[rank0]: ^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1830, in inner
[rank0]: result = forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/utils/generic.py", line 903, in wrapper
[rank0]: output = func(self, *args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py", line 2015, in forward
[rank0]: outputs = self.model(
[rank0]: ^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1776, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1882, in _call_impl
[rank0]: return inner()
[rank0]: ^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1830, in inner
[rank0]: result = forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/utils/generic.py", line 903, in wrapper
[rank0]: output = func(self, *args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py", line 1703, in forward
[rank0]: outputs = self.language_model(
[rank0]: ^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1776, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1882, in _call_impl
[rank0]: return inner()
[rank0]: ^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1830, in inner
[rank0]: result = forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/utils/generic.py", line 1032, in wrapper
[rank0]: output = func(self, *args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/utils/output_capturing.py", line 252, in wrapper
[rank0]: outputs = func(self, *args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py", line 1308, in forward
[rank0]: hidden_states = decoder_layer(
[rank0]: ^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/modeling_layers.py", line 92, in __call__
[rank0]: return self._gradient_checkpointing_func(partial(super().__call__, **kwargs), *args)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/juicefs-data/zhaoyiru/llm_train/ms-swift/swift/trainers/mixin.py", line 846, in _new_checkpoint
[rank0]: return _old_checkpoint(*args, use_reentrant=use_reentrant_, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/_compile.py", line 54, in inner
[rank0]: return disable_fn(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py", line 1181, in _fn
[rank0]: return fn(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/utils/checkpoint.py", line 505, in checkpoint
[rank0]: return CheckpointFunction.apply(function, preserve, *args)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/autograd/function.py", line 583, in apply
[rank0]: return super().apply(*args, **kwargs) # type: ignore[misc]
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/utils/checkpoint.py", line 268, in forward
[rank0]: outputs = run_function(*args)
[rank0]: ^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1776, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1882, in _call_impl
[rank0]: return inner()
[rank0]: ^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1830, in inner
[rank0]: result = forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py", line 865, in forward
[rank0]: hidden_states = self.linear_attn(
[rank0]: ^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1776, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1882, in _call_impl
[rank0]: return inner()
[rank0]: ^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1830, in inner
[rank0]: result = forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py", line 532, in forward
[rank0]: core_attn_out, last_recurrent_state = self.chunk_gated_delta_rule(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py", line 1181, in _fn
[rank0]: return fn(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/fla/ops/gated_delta_rule/chunk.py", line 405, in chunk_gated_delta_rule
[rank0]: o, final_state = ChunkGatedDeltaRuleFunction.apply(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/autograd/function.py", line 583, in apply
[rank0]: return super().apply(*args, **kwargs) # type: ignore[misc]
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/fla/utils.py", line 222, in wrapper
[rank0]: return fn(*processed_args, **processed_kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/amp/autocast_mode.py", line 477, in decorate_fwd
[rank0]: return fwd(*args, **kwargs) # pyrefly: ignore [not-callable]
[rank0]: ^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/fla/ops/gated_delta_rule/chunk.py", line 247, in forward
[rank0]: g, o, A, final_state, initial_state = chunk_gated_delta_rule_fwd(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/fla/ops/gated_delta_rule/chunk.py", line 40, in chunk_gated_delta_rule_fwd
[rank0]: A = chunk_scaled_dot_kkt_fwd(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/fla/ops/common/chunk_scaled_dot_kkt.py", line 114, in chunk_scaled_dot_kkt_fwd
[rank0]: chunk_scaled_dot_kkt_fwd_kernel[(NT, B * H)](
[rank0]: File "/usr/local/lib/python3.12/site-packages/triton/runtime/jit.py", line 390, in <lambda>
[rank0]: return lambda *args, **kwargs: self.run(grid=grid, warmup=False, *args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/triton/runtime/autotuner.py", line 453, in run
[rank0]: return self.fn.run(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/triton/runtime/autotuner.py", line 237, in run
[rank0]: used_cached_result = self.check_disk_cache(key, pruned_configs, benchmark)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/triton/runtime/autotuner.py", line 201, in check_disk_cache
[rank0]: bench_fn()
[rank0]: File "/usr/local/lib/python3.12/site-packages/triton/runtime/autotuner.py", line 228, in benchmark
[rank0]: timings = {config: self._bench(*args, config=config, **kwargs) for config in pruned_configs}
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/triton/runtime/autotuner.py", line 160, in _bench
[rank0]: return self.do_bench(kernel_call, quantiles=(0.5, 0.2, 0.8))
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/site-packages/triton/testing.py", line 150, in do_bench
[rank0]: di.synchronize()
[rank0]: File "/usr/local/lib/python3.12/site-packages/torch/cuda/__init__.py", line 1108, in synchronize
[rank0]: return torch._C._cuda_synchronize()
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: torch.AcceleratorError: CUDA error: an illegal memory access was encountered
[rank0]: Search for `cudaErrorIllegalAddress' in https://docs.nvidia.com/cuda/cuda-runtime-api/group__CUDART__TYPES.html for more information.
[rank0]: CUDA kernel errors might be asynchronously reported at some other API call, so the stacktrace below might be incorrect.
[rank0]: For debugging consider passing CUDA_LAUNCH_BLOCKING=1
[rank0]: Compile with `TORCH_USE_CUDA_DSA` to enable device-side assertions.
How to Reproduce / 如何复现
训练脚本:
NPROC_PER_NODE=8 \
CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7 \
PYTORCH_ALLOC_CONF=expandable_segments:True \
swift rlhf \
--rlhf_type dpo \
--model models/Qwen3.6-35B-A3B \
--tuner_type full \
--agent_template qwen3_coder \
--enable_thinking true \
--dataset xxx.jsonl \
--load_from_cache_file false \
--split_dataset_ratio 0.01 \
--torch_dtype bfloat16 \
--num_train_epochs 1 \
--per_device_train_batch_size 1 \
--per_device_eval_batch_size 1 \
--learning_rate 1e-5 \
--gradient_accumulation_steps 2 \
--eval_steps 100 \
--save_steps 100 \
--save_total_limit 2 \
--logging_steps 5 \
--max_length 32768 \
--output_dir checkpoints/dpo \
--warmup_ratio 0.05 \
--save_only_model true \
--dataloader_num_workers 4 \
--dataset_num_proc 4 \
--deepspeed zero3_offload \
--attn_impl flash_attn \
--rpo_alpha 0.1 \
--padding_free true \
--truncation_strategy left \
--sequence_parallel_size 4
环境版本:
transformers 5.10.2
ms_swift 4.4.0.dev0
fla-core 0.4.2
flash-linear-attention 0.4.2
triton 3.4.0
Additional Information / 补充信息
No response
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by reproducing DPO training with --sequence_parallel_size 4, comparing the models reported in the issue. Trace from swift/rlhf_trainers/dpo_trainer.py, especially concatenated_forward and get_batch_loss_metrics, into modeling_qwen3_5_moe.py and fla/ops/gated_delta_rule/chunk.py. Done means the reported configuration runs without the CUDA illegal-memory-access error.
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