Skip to content

Commit 34a6c0a

Browse files
jhvmhgpre-commit-ci[bot]sudhakarsingh27
authored
Fix Flash Attention 3 API compatibility for window size parameters (NVIDIA#2704)
* Fix Flash Attention 3 API compatibility for window size parameters Replace single window_size parameter with window_size_left and window_size_right in flash_attn_fwd function to align with flash-attn v2.7.0+ API changes. - Update function signature in flash_attn_interface - Maintain backward compatibility where possible - Ensure consistency with Flash Attention v2 implementation Signed-off-by: Chaoyang Mei <1192554423@qq.com> Signed-off-by: meichaoyang001 <meichaoyang001@ke.com> * Fix Flash Attention 3 backward API parameter naming Rename causal parameter to is_causal in flash_attn_bwd function to align with flash-attn v2.7.0+ API changes. This ensures consistency with the updated flash-attn library interface for backward pass operations. Signed-off-by: meichaoyang001 <meichaoyang001@ke.com> * Fix Flash Attention 3 backward API parameter naming Rename causal parameter to is_causal in flash_attn_bwd function to align with flash-attn v2.7.0+ API changes. This ensures consistency with the updated flash-attn library interface for backward pass operations. Signed-off-by: meichaoyang001 <meichaoyang001@ke.com> * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Refactor Flash Attention 3 to use positional args instead of kwargs Replace keyword arguments with positional arguments in flash_attn_fwd and flash_attn_bwd to abstract away parameter naming differences (causal vs is_causal) between flash-attn versions. This provides a more robust interface that is resilient to future API changes in the flash-attn library. - Convert window_size_left, window_size_right, and causal parameters to positional args in both forward and backward functions - Eliminate version-specific parameter naming dependencies - Simplify compatibility handling across flash-attn v2.7.0+ variants Signed-off-by: meichaoyang001 <meichaoyang001@ke.com> * Fix Flash Attention 3 backward API parameter naming Rename causal parameter to is_causal in flash_attn_bwd function to align with flash-attn v3 API changes. This ensures consistency with the updated flash-attn library interface for backward pass operations. Signed-off-by: meichaoyang001 <meichaoyang001@ke.com> --------- Signed-off-by: Chaoyang Mei <1192554423@qq.com> Signed-off-by: meichaoyang001 <meichaoyang001@ke.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Sudhakar Singh <sudhakars@nvidia.com>
1 parent 6638fef commit 34a6c0a

1 file changed

Lines changed: 28 additions & 18 deletions

File tree

transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py

Lines changed: 28 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -937,9 +937,9 @@ def cp_p2p_fwd_flash_attn(
937937
elif section == "upper-triangle":
938938
max_seqlen_q_ = max_seqlen_q // 2
939939
if section in ["lower-triangle", "upper-triangle"]:
940-
if use_flash_attn_3 or (fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus):
940+
if fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus:
941941
fa_forward_kwargs["window_size"] = (-1, -1)
942-
elif fa_utils.v2_7_0_plus:
942+
elif use_flash_attn_3 or fa_utils.v2_7_0_plus:
943943
fa_forward_kwargs["window_size_left"] = -1
944944
fa_forward_kwargs["window_size_right"] = -1
945945

@@ -1189,9 +1189,9 @@ def cp_p2p_bwd_flash_attn(
11891189
):
11901190
"""Per-tile backward call of CP P2P with FlashAttention backend"""
11911191
dq, dk, dv = [torch.empty_like(x) for x in [q_part, k_part, v_part]]
1192-
if use_flash_attn_3 or (fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus):
1192+
if fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus:
11931193
fa_backward_kwargs["window_size"] = (-1, -1)
1194-
elif fa_utils.v2_7_0_plus:
1194+
elif use_flash_attn_3 or fa_utils.v2_7_0_plus:
11951195
fa_backward_kwargs["window_size_left"] = -1
11961196
fa_backward_kwargs["window_size_right"] = -1
11971197
if not use_flash_attn_3:
@@ -1201,9 +1201,9 @@ def cp_p2p_bwd_flash_attn(
12011201
softmax_lse__ = softmax_lse
12021202
causal_ = False
12031203
if section == "diagonal":
1204-
if use_flash_attn_3 or (fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus):
1204+
if fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus:
12051205
fa_backward_kwargs["window_size"] = (-1, 0)
1206-
elif fa_utils.v2_7_0_plus:
1206+
elif use_flash_attn_3 or fa_utils.v2_7_0_plus:
12071207
fa_backward_kwargs["window_size_left"] = -1
12081208
fa_backward_kwargs["window_size_right"] = 0
12091209
causal_ = True
@@ -1225,6 +1225,10 @@ def cp_p2p_bwd_flash_attn(
12251225
dk=dk,
12261226
dv=dv,
12271227
)
1228+
if use_flash_attn_3:
1229+
fa_backward_kwargs["is_causal"] = causal_
1230+
else:
1231+
fa_backward_kwargs["causal"] = causal_
12281232
flash_attn_bwd(
12291233
dout_part,
12301234
q_part,
@@ -1233,7 +1237,6 @@ def cp_p2p_bwd_flash_attn(
12331237
out_part,
12341238
softmax_lse__,
12351239
*fa_backward_args_thd,
1236-
causal=causal_,
12371240
**fa_backward_kwargs,
12381241
)
12391242

@@ -1508,7 +1511,8 @@ def forward(
15081511
flash_attn_fwd = (
15091512
_flash_attn_fwd_v3 # pylint: disable=possibly-used-before-assignment
15101513
)
1511-
fa_forward_kwargs["window_size"] = (-1, 0) if causal else (-1, -1)
1514+
fa_forward_kwargs["window_size_left"] = -1
1515+
fa_forward_kwargs["window_size_right"] = 0 if causal else -1
15121516
else:
15131517
if qkv_format == "thd":
15141518
from transformer_engine.pytorch.attention.dot_product_attention.backends import (
@@ -2985,9 +2989,9 @@ def forward(
29852989
max_seqlen_q=max_seqlen_q,
29862990
max_seqlen_kv=max_seqlen_kv_,
29872991
)
2988-
if use_flash_attn_3 or (fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus):
2992+
if fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus:
29892993
fa_forward_kwargs["window_size"] = window_size_per_step[i]
2990-
elif fa_utils.v2_7_0_plus:
2994+
elif use_flash_attn_3 or fa_utils.v2_7_0_plus:
29912995
fa_forward_kwargs["window_size_left"] = window_size_per_step[i][0]
29922996
fa_forward_kwargs["window_size_right"] = window_size_per_step[i][1]
29932997
fa_outputs = flash_attn_fwd(
@@ -3206,13 +3210,15 @@ def backward(ctx, dout, *_args):
32063210
)
32073211
if not ctx.use_flash_attn_3:
32083212
fa_backward_kwargs["rng_state"] = rng_states[i]
3209-
if ctx.use_flash_attn_3 or (
3210-
fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus
3211-
):
3213+
if fa_utils.v2_3_plus and not fa_utils.v2_7_0_plus:
32123214
fa_backward_kwargs["window_size"] = window_size_per_step[i]
3213-
elif fa_utils.v2_7_0_plus:
3215+
elif ctx.use_flash_attn_3 or fa_utils.v2_7_0_plus:
32143216
fa_backward_kwargs["window_size_left"] = window_size_per_step[i][0]
32153217
fa_backward_kwargs["window_size_right"] = window_size_per_step[i][1]
3218+
if ctx.use_flash_attn_3:
3219+
fa_backward_kwargs["is_causal"] = "causal" in ctx.attn_mask_type
3220+
else:
3221+
fa_backward_kwargs["causal"] = "causal" in ctx.attn_mask_type
32163222
flash_attn_bwd(
32173223
dout_,
32183224
q_,
@@ -3221,7 +3227,6 @@ def backward(ctx, dout, *_args):
32213227
out_,
32223228
softmax_lse_per_step[i],
32233229
*fa_backward_args_thd,
3224-
causal="causal" in ctx.attn_mask_type,
32253230
**fa_backward_kwargs,
32263231
)
32273232

@@ -3361,7 +3366,8 @@ def forward(
33613366
)
33623367

33633368
flash_attn_fwd = _flash_attn_fwd_v3
3364-
fa_forward_kwargs["window_size"] = window_size
3369+
fa_forward_kwargs["window_size_left"] = window_size[0]
3370+
fa_forward_kwargs["window_size_right"] = window_size[1]
33653371
else:
33663372
if qkv_format == "thd":
33673373
from transformer_engine.pytorch.attention.dot_product_attention.backends import (
@@ -3738,7 +3744,8 @@ def backward(ctx, dout, *_args):
37383744
flash_attn_bwd = (
37393745
_flash_attn_bwd_v3 # pylint: disable=possibly-used-before-assignment
37403746
)
3741-
fa_backward_kwargs["window_size"] = ctx.window_size
3747+
fa_backward_kwargs["window_size_left"] = ctx.window_size[0]
3748+
fa_backward_kwargs["window_size_right"] = ctx.window_size[1]
37423749
fa_backward_kwargs["deterministic"] = ctx.deterministic
37433750
else:
37443751
if qkv_format == "thd":
@@ -3821,6 +3828,10 @@ def backward(ctx, dout, *_args):
38213828
)
38223829
if not ctx.use_flash_attn_3:
38233830
fa_backward_kwargs["rng_state"] = rng_state
3831+
fa_backward_kwargs["causal"] = causal
3832+
else:
3833+
fa_backward_kwargs["is_causal"] = causal
3834+
38243835
flash_attn_bwd(
38253836
dout,
38263837
q,
@@ -3829,7 +3840,6 @@ def backward(ctx, dout, *_args):
38293840
out,
38303841
softmax_lse,
38313842
*fa_backward_args_thd,
3832-
causal=causal,
38333843
**fa_backward_kwargs,
38343844
)
38353845

0 commit comments

Comments
 (0)