diff --git a/src/paddlefleet/context_parallel_utils.py b/src/paddlefleet/context_parallel_utils.py index 3bd9fe3e5d..41df7c9a35 100644 --- a/src/paddlefleet/context_parallel_utils.py +++ b/src/paddlefleet/context_parallel_utils.py @@ -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: + 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( diff --git a/src/paddlefleet/refined_recompute/flash_attn.py b/src/paddlefleet/refined_recompute/flash_attn.py index 7e7f605301..6dfa6bf3ba 100644 --- a/src/paddlefleet/refined_recompute/flash_attn.py +++ b/src/paddlefleet/refined_recompute/flash_attn.py @@ -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: + 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(