Files
paddlepaddle--paddle/test/test_flashmask_ci/test_flashmask_ci.py
T
2026-07-13 12:40:42 +08:00

453 lines
15 KiB
Python

# 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,
)