Skip to content

feat(sp): support flash_attention_3 in sequence parallelism - #294

Open
werwrewe wants to merge 2 commits into
modelscope:dev-swift-v5from
werwrewe:dev-swift-v5-fa3
Open

werwrewe wants to merge 2 commits into
modelscope:dev-swift-v5from
werwrewe:dev-swift-v5-fa3

Conversation

@werwrewe

Copy link
Copy Markdown
Contributor

PR type

  • Bug Fix
  • New Feature
  • Document Updates
  • More Models or Datasets Support

PR information

Add FlashAttention-3 support for CUDA and ring attn.

Experiment results

1a. Qwen3-0.6B × transformers 5.15.1

拓扑 后端 batch 数据长度 loss(首→末) 耗时
sp=3(ring) fa3 1 ~4096 2.2668→1.6629 88s
sp=3(ring) fa2 1 ~4096 2.2675→1.6629 92s
sp=2(ulysses) fa3 8 ~4096 2.2191→1.6927 61s
sp=2(ulysses) fa2 8 ~4096 2.2194→1.6927 63s

1b. Qwen3.5-4B × transformers 5.15.1

拓扑 后端 batch 数据长度 loss(首→末) 耗时
sp=3(ring) fa3 — ~4096 不支持 ❌ —
sp=3(ring) fa2 — ~4096 不支持 ❌ —
sp=2(ulysses) fa3 4 ~4096 1.9881→1.5512 258s
sp=2(ulysses) fa2 4 ~4096 1.9888→1.5496 259s

1c. Qwen3-0.6B × transformers 4.57.0

拓扑 后端 batch 数据长度 loss(首→末) 耗时
sp=2(ulysses) fa3 8 ~4096 ⚠️ 2.2191→5.55→4.32 58s
sp=2(ulysses) fa2 8 ~4096 ⚠️ 2.2194→5.55→4.29 61s
sp=3(ring) fa3 1 ~4096 2.2668→1.6630 78s
sp=3(ring) fa2 1 ~4096 2.2675→1.6628 72s

fa2/fa3 share transformers' single flash_attention_forward entry and
dispatch to the concrete kernel at call time via
config._attn_implementation. Register the SP wrapper under the fa3 key
as well (exact-key registry lookup would otherwise silently train
without SP), and add the fa3 leaf to zigzag ring attention via
flash_attn_interface (gradients through dq/dk/dv buffer mutation).
Also log the requested vs resolved attn_implementation at model build.

Verified on H800: fa3/fa2 training curves match to bf16 noise on
Qwen3-0.6B (ring, ulysses) and Qwen3.5-4B (ulysses), transformers
4.57.0 and 5.15.1.
Companion to dadb3f0 which moved the helper out of vllm_sampler_tq;
the stale import broke test collection in CI.
local_flash_attn, dist_attn=DistributedAttention(None, self))
for _impl in FLASH_ATTENTION_IMPLS:
ALL_ATTENTION_FUNCTIONS[_impl] = partial(
local_flash_attn, dist_attn=DistributedAttention(None, self))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

有 dropout 但是用了fa3 就报错吧;还有就是做一下精度对齐,前向loss和反向梯度都和sdpa 比一下

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

精度对齐(H800,fa3 vs fa2,单卡 varlen,H=8 D=128,LENS=[100,300,512],transformers=5.15.1):

检查项 结果
out 数值 max|Δ|=3.91e-03
lse 数值 max|Δ|=4.77e-07
dq 数值 max|Δ|=3.91e-03
dk 数值 max|Δ|=7.81e-03
dv 数值 max|Δ|=3.91e-03

unit test:

测试 结果
sp_only_multi_sample PASS
sp_only_multi_sample_masked PASS
cp_only PASS
cp_sp PASS

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants