# Copyright (c) 2025 PaddlePaddle Authors. All Rights Reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. from functools import partial import pytest from generate_startend_row_indices import ( generate_causal_blockwise_mask, generate_causal_document_mask, generate_document_mask, generate_global_sliding_window_mask, generate_none_mask, generate_prefix_lm_causal_mask, generate_prefix_lm_document_mask, generate_qk_sparse_mask, generate_random_eviction_mask, generate_share_question_mask, generate_sliding_window_mask, startend_row_indices_to_attn_bias, ) from test_util import attention_ref import paddle from paddle.nn.functional.flash_attention import flashmask_attention # batch_size, seqlen_q, seqlen_k, nheads, nheads_kv shape_cases = [ (2840, 32, 32, 16, 4), (1, 300, 300, 16, 16), # (2, 8192, 32768, 32, 4), # this will oom # (2, 8192, 8192, 32, 4), # this will oom (2, 8192, 8192, 14, 1), (2, 16384, 16384, 4, 1), (1, 1, 127, 1, 1), (1, 128, 127, 1, 1), (1, 127, 128, 1, 1), (2, 16383, 16384, 4, 1), (2, 16384, 16383, 4, 1), (2, 1000, 1000, 4, 1), (2, 2000, 2000, 4, 1), (2, 3000, 3000, 4, 1), (1, 4000, 4000, 1, 1), (1, 8192, 32768 + 1024, 2, 1), (1, 8192, 16384 + 1024, 2, 1), # my case ] # Generate all combinations for second param def generate_shapes(): for batch_size, seqlen_q, seqlen_k, nheads, nheads_kv in shape_cases: if nheads_kv == 1: nheads_startend_row_indices_values = [1] else: nheads_startend_row_indices_values = [1, nheads_kv] for nheads_startend_row_indices in nheads_startend_row_indices_values: yield ( batch_size, seqlen_q, seqlen_k, nheads, nheads_kv, nheads_startend_row_indices, ) @pytest.mark.parametrize("dtype", [paddle.bfloat16]) @pytest.mark.parametrize("fa_version", [3]) @pytest.mark.parametrize("d, dv", [(128, 128), (80, 80), (64, 64), (256, 256)]) @pytest.mark.parametrize( "batch_size, seqlen_q, seqlen_k, nheads, nheads_kv, nheads_startend_row_indices", list(generate_shapes()), ) @pytest.mark.parametrize( "gen_startend_row_indices", [ partial(generate_none_mask, causal=False), # full partial(generate_none_mask, causal=True), # causal partial(generate_sliding_window_mask), # sliding window partial(generate_causal_document_mask), # causal document mask partial(generate_document_mask), # document mask partial(generate_share_question_mask), # share question mask partial(generate_global_sliding_window_mask), # global sliding window partial(generate_causal_blockwise_mask), # causal blockwise mask partial(generate_prefix_lm_document_mask), # prefix lm document mask partial(generate_prefix_lm_causal_mask), # prefix lm causal mask partial(generate_qk_sparse_mask), # qk-sparse mask partial(generate_random_eviction_mask), # random eviction mask ], ) def test_flashmask( batch_size, seqlen_q, seqlen_k, nheads, nheads_kv, d, dv, nheads_startend_row_indices, fa_version, dtype, gen_startend_row_indices, softcap=0.0, ): paddle.seed(2024) assert nheads % nheads_kv == 0 q_ref = paddle.randn(shape=[batch_size, seqlen_q, nheads, d], dtype=dtype) k_ref = paddle.randn( shape=[batch_size, seqlen_k, nheads_kv, d], dtype=dtype ) v_ref = paddle.randn( shape=[batch_size, seqlen_k, nheads_kv, dv], dtype=dtype ) q_ref.stop_gradient = False k_ref.stop_gradient = False v_ref.stop_gradient = False q_bf16, k_bf16, v_bf16 = [x.detach().clone() for x in (q_ref, k_ref, v_ref)] q_bf16.stop_gradient = False k_bf16.stop_gradient = False v_bf16.stop_gradient = False q, k, v = [x.detach().clone() for x in (q_ref, k_ref, v_ref)] q.stop_gradient = False k.stop_gradient = False v.stop_gradient = False startend_row_indices, causal = gen_startend_row_indices( batch_size, seqlen_q, seqlen_k, nheads_startend_row_indices ) if startend_row_indices is None and causal and d == 80: pytest.skip( "Skipping because running headdim 80 with flash_attn in causal mask" ) attn_bias = startend_row_indices_to_attn_bias( startend_row_indices, seqlen_q, nheads, dtype, causal ) out_ref, attn_ref = attention_ref( q_ref, k_ref, v_ref, causal=causal, attn_bias=attn_bias ) out_bf16, attn_bf16 = attention_ref( q_bf16, k_bf16, v_bf16, causal=causal, attn_bias=attn_bias, upcast=False, reorder_ops=True, ) # # Numerical error if we just do any arithmetic on out_ref fwd_atol = 2 * (out_ref + 0.3 - 0.3 - out_ref).abs().max().item() assert softcap == 0.0 rtol = 2 if softcap == 0.0 else 3 print( f"Paddle naive bf16 Output max diff: {(out_bf16 - out_ref).abs().max().item()}" ) print( f"Paddle naive bf16 Output mean diff: {(out_bf16 - out_ref).abs().mean().item()}" ) if fa_version == 2: paddle.set_flags({'FLAGS_flash_attn_version': 2}) elif fa_version == 3: paddle.set_flags({'FLAGS_flash_attn_version': 3}) else: raise ValueError(f"Invalid flash attention version: {fa_version}") out, lse = flashmask_attention( q, k, v, startend_row_indices=startend_row_indices, causal=causal, return_softmax_lse=True, ) print(f"flashmask Output max diff: {(out - out_ref).abs().max().item()}") print(f"flashmask Output mean diff: {(out - out_ref).abs().mean().item()}") # if not causal: # print(f"LSE max diff: {(lse - lse_ref).abs().max().item()}") # breakpoint() # Check that FlashAttention's numerical error is at most twice the numerical error # of a Pytorch implementation. assert (out - out_ref).abs().max().item() <= rtol * ( out_bf16 - out_ref ).abs().max().item() + fwd_atol g = paddle.randn(shape=out.shape, dtype=out.dtype) out.backward(g) out_ref.backward(g) out_bf16.backward(g) print(f"flashmask dQ max diff: {(q.grad - q_ref.grad).abs().max().item()}") print(f"flashmask dK max diff: {(k.grad - k_ref.grad).abs().max().item()}") print(f"flashmask dV max diff: {(v.grad - v_ref.grad).abs().max().item()}") print( f"flashmask dQ mean diff: {(q.grad - q_ref.grad).abs().mean().item()}" ) print( f"flashmask dK mean diff: {(k.grad - k_ref.grad).abs().mean().item()}" ) print( f"flashmask dV mean diff: {(v.grad - v_ref.grad).abs().mean().item()}" ) print( f"Paddle naive bf16 dQ max diff: {(q_bf16.grad - q_ref.grad).abs().max().item()}" ) print( f"Paddle naive bf16 dK max diff: {(k_bf16.grad - k_ref.grad).abs().max().item()}" ) print( f"Paddle naive bf16 dV max diff: {(v_bf16.grad - v_ref.grad).abs().max().item()}" ) print( f"Paddle naive bf16 dQ mean diff: {(q_bf16.grad - q_ref.grad).abs().mean().item()}" ) print( f"Paddle naive bf16 dK mean diff: {(k_bf16.grad - k_ref.grad).abs().mean().item()}" ) print( f"Paddle naive bf16 dV mean diff: {(v_bf16.grad - v_ref.grad).abs().mean().item()}" ) dq_atol = 2 * (q_ref.grad + 0.3 - 0.3 - q_ref.grad).abs().max().item() + ( 0 if softcap == 0 else 3e-4 ) assert (q.grad - q_ref.grad).abs().max().item() <= rtol * ( q_bf16.grad - q_ref.grad ).abs().max().item() + dq_atol dk_atol = 2 * (k_ref.grad + 0.3 - 0.3 - k_ref.grad).abs().max().item() + ( 0 if softcap == 0 else 3e-4 ) assert (k.grad - k_ref.grad).abs().max().item() <= rtol * ( k_bf16.grad - k_ref.grad ).abs().max().item() + dk_atol dv_atol = 2 * (v_ref.grad + 0.3 - 0.3 - v_ref.grad).abs().max().item() + ( 0 if softcap == 0 else 3e-4 ) assert (v.grad - v_ref.grad).abs().max().item() <= rtol * ( v_bf16.grad - v_ref.grad ).abs().max().item() + dv_atol @pytest.mark.parametrize("d, dv", [(192, 128), (256, 128)]) @pytest.mark.parametrize("fa_version", [2, 3]) @pytest.mark.parametrize( "gen_startend_row_indices", [ partial(generate_none_mask, causal=False), # flash_attn path partial(generate_sliding_window_mask), # flashmask path ], ) def test_flashmask_forward_headdim_mismatch_raises( d, dv, fa_version, gen_startend_row_indices ): """Test that headdim != headdim_v raises an error.""" if fa_version == 2: paddle.set_flags({'FLAGS_flash_attn_version': 2}) elif fa_version == 3: paddle.set_flags({'FLAGS_flash_attn_version': 3}) else: raise ValueError(f"Invalid flash attention version: {fa_version}") batch_size, seqlen_q, seqlen_k, nheads, nheads_startend_row_indices = ( 1, 300, 300, 16, 16, ) q = paddle.randn([batch_size, seqlen_q, nheads, d], dtype=paddle.bfloat16) k = paddle.randn([batch_size, seqlen_k, nheads, d], dtype=paddle.bfloat16) v = paddle.randn([batch_size, seqlen_k, nheads, dv], dtype=paddle.bfloat16) startend_row_indices, causal = gen_startend_row_indices( batch_size, seqlen_q, seqlen_k, nheads_startend_row_indices ) if fa_version == 3: if startend_row_indices is None: # fallback to fa2 with pytest.raises(Exception, match="headdim != headdim_v"): out, lse = flashmask_attention( q, k, v, startend_row_indices=startend_row_indices, causal=causal, return_softmax_lse=True, ) else: # flashmask v3 if not ( (d > 128 and d <= 192 and dv > 96 and dv <= 128) or (d <= 64 and dv <= 512) ): with pytest.raises( Exception, match="headdim != headdim_v|V headdim is different from", ): out, lse = flashmask_attention( q, k, v, startend_row_indices=startend_row_indices, causal=causal, return_softmax_lse=True, ) elif fa_version == 2: with pytest.raises(Exception, match="headdim != headdim_v"): out, lse = flashmask_attention( q, k, v, startend_row_indices=startend_row_indices, causal=causal, return_softmax_lse=True, ) else: raise ValueError(f"Invalid flash attention version: {fa_version}") @pytest.mark.parametrize("d, dv", [(192, 128), (256, 128)]) @pytest.mark.parametrize("fa_version", [2, 3]) def test_flashmask_backward_headdim_mismatch_raises(fa_version, d, dv): """Test that backward kernel raises when headdim != headdim_v.""" from paddle import _C_ops paddle.set_flags({'FLAGS_flash_attn_version': fa_version}) batch_size, seqlen_q, seqlen_k, nheads = 1, 300, 300, 16 softmax_scale = d ** (-0.5) if fa_version == 2: q = paddle.randn( [batch_size, seqlen_q, nheads, d], dtype=paddle.bfloat16 ) k = paddle.randn( [batch_size, seqlen_k, nheads, d], dtype=paddle.bfloat16 ) v = paddle.randn( [batch_size, seqlen_k, nheads, dv], dtype=paddle.bfloat16 ) out = paddle.randn( [batch_size, seqlen_q, nheads, dv], dtype=paddle.bfloat16 ) softmax_lse = paddle.randn( [batch_size, nheads, seqlen_q], dtype=paddle.float32 ) seed_offset = paddle.to_tensor([0, 0], dtype=paddle.int64) dout = paddle.randn( [batch_size, seqlen_q, nheads, dv], dtype=paddle.bfloat16 ) with pytest.raises(Exception, match="headdim != headdim_v"): _C_ops.flash_attn_grad( q, k, v, out, softmax_lse, seed_offset, None, dout, 0.0, False ) elif fa_version == 3: total_q = batch_size * seqlen_q total_k = batch_size * seqlen_k q = paddle.randn([total_q, nheads, d], dtype=paddle.bfloat16) k = paddle.randn([total_k, nheads, d], dtype=paddle.bfloat16) v = paddle.randn([total_k, nheads, dv], dtype=paddle.bfloat16) out = paddle.randn([total_q, nheads, dv], dtype=paddle.bfloat16) softmax_lse = paddle.randn( [batch_size, nheads, seqlen_q], dtype=paddle.float32 ) cu_seqlens_q = paddle.to_tensor([0, seqlen_q], dtype=paddle.int32) cu_seqlens_k = paddle.to_tensor([0, seqlen_k], dtype=paddle.int32) dout = paddle.randn([total_q, nheads, dv], dtype=paddle.bfloat16) with pytest.raises(Exception, match="headdim != headdim_v"): _C_ops.flash_attn_v3_varlen_grad( q, k, v, out, softmax_lse, cu_seqlens_q, cu_seqlens_k, None, None, # seqused_q, seqused_k dout, softmax_scale, seqlen_q, seqlen_k, False, # causal -1, -1, # window_size_left, window_size_right 0.0, # softcap 0, # sm_margin ) @pytest.mark.parametrize("d, dv", [(192, 128), (256, 128)]) def test_flash_attn_unpadded_grad_headdim_mismatch_raises(d, dv): """Test that flash_attn_unpadded_grad raises when headdim != headdim_v.""" from paddle import _C_ops paddle.set_flags({'FLAGS_flash_attn_version': 2}) batch_size, seqlen_q, seqlen_k, nheads = 1, 300, 300, 16 total_q = batch_size * seqlen_q total_k = batch_size * seqlen_k q = paddle.randn([total_q, nheads, d], dtype=paddle.bfloat16) k = paddle.randn([total_k, nheads, d], dtype=paddle.bfloat16) v = paddle.randn([total_k, nheads, dv], dtype=paddle.bfloat16) out = paddle.randn([total_q, nheads, dv], dtype=paddle.bfloat16) softmax_lse = paddle.randn( [batch_size, nheads, seqlen_q], dtype=paddle.float32 ) seed_offset = paddle.to_tensor([0, 0], dtype=paddle.int64) cu_seqlens_q = paddle.to_tensor([0, seqlen_q], dtype=paddle.int32) cu_seqlens_k = paddle.to_tensor([0, seqlen_k], dtype=paddle.int32) dout = paddle.randn([total_q, nheads, dv], dtype=paddle.bfloat16) with pytest.raises(Exception, match="headdim != headdim_v"): _C_ops.flash_attn_unpadded_grad( q, k, v, cu_seqlens_q, cu_seqlens_k, out, softmax_lse, seed_offset, None, dout, seqlen_q, seqlen_k, d ** (-0.5), 0.0, False, )