Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions src/paddlefleet/context_parallel_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -730,11 +730,12 @@ def cp_flashmask_allgatherkv_balance_forward(
deterministic = paddle.get_flags(["FLAGS_cudnn_deterministic"])[
"FLAGS_cudnn_deterministic"
]
if "block_mask" in inspect.signature(flashmask_attention).parameters:
if deterministic and query.shape[-1] > 128:
if fa_version == 3:

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

🔴 Bug CP 路径同样只检查原始 fa_version == 3,导致请求 v4 但 _flash_mask_available=False 时前向实际走普通 flashmask_attention,却仍把 fa_version=4 返回给后向。

cp_flashmask_allgatherkv_balance_backward() 只有在 fa_version == 4 and _flash_mask_available 时才走 v4 分支,否则会落到 ValueError("FlashAttention version 4 is not supported.")。确定性配置下旧逻辑会降到 v2,新逻辑跳过降级后会让 CP 训练在 backward 阶段失败。

建议修复方式:
_get_fa_version() 一样集中计算 effective version;当 v4 后端不可用时把 fa_version 改成普通 FlashMask 实际使用的 3/2,并把这个值返回给 backward。

if "block_mask" in inspect.signature(flashmask_attention).parameters:
if deterministic and query.shape[-1] > 128:
fa_version = 2
elif deterministic:
fa_version = 2
elif deterministic:
fa_version = 2

if fa_version == 4 and _flash_mask_available:
output, log_sum_exp = _flash_attn_fwd(
Expand Down
25 changes: 13 additions & 12 deletions src/paddlefleet/refined_recompute/flash_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,21 +61,22 @@ def _get_fa_version(hdim):
if "xpu" in paddle.get_device():
return 2
# Xiangrui: For deterministic, NOT support for hdim > 128 currently.
if "block_mask" in inspect.signature(flashmask_attention).parameters:
if (
paddle.get_flags(["FLAGS_cudnn_deterministic"])[
"FLAGS_cudnn_deterministic"
]
and hdim > 128
):
return 2
elif paddle.get_flags(["FLAGS_cudnn_deterministic"])[
"FLAGS_cudnn_deterministic"
]:
return 2
fa_version = paddle.base.framework.get_flags(["FLAGS_flash_attn_version"])[
"FLAGS_flash_attn_version"
]
if fa_version == 3:

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

🔴 Bug 这里用原始 flag 是否等于 3 来决定 deterministic fallback,会漏掉 FLAGS_flash_attn_version=4_flash_mask_available=False 的回退路径。

当前函数后面会把不可用的 v4 回退成 v3;但这个回退发生在 deterministic 检查之后。确定性模式且 hdim > 128(或旧签名无 block_mask)时,原逻辑会返回 v2,现在会返回 v3,违背上面的 deterministic 不支持约束,并可能在 recompute 的前/后向中调用不支持的 v3 kernel。

建议修复方式:
先把用户请求的版本归一化为有效版本,再对有效的 v3 执行 deterministic fallback,例如:

if fa_version == 4 and not _flash_mask_available:
    fa_version = 3
if fa_version == 3 and deterministic_needs_v2:
    return 2
return fa_version

if "block_mask" in inspect.signature(flashmask_attention).parameters:
if (
paddle.get_flags(["FLAGS_cudnn_deterministic"])[
"FLAGS_cudnn_deterministic"
]
and hdim > 128
):
return 2
elif paddle.get_flags(["FLAGS_cudnn_deterministic"])[
"FLAGS_cudnn_deterministic"
]:
return 2
# Fall back to version 3 if flash_mask is not available
if fa_version == 4 and not _flash_mask_available:
logger.warning(
Expand Down
Loading