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

136 lines
4.9 KiB
Python

# Copyright (c) 2024 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.
# Unit test for paddle.nn.functional.flash_attention
# Target: cover _select_sdp_cuda, _select_sdp, flash_attention, scaled_dot_product_attention
import unittest
import paddle
from paddle.nn.functional.flash_attention import (
_select_sdp,
_select_sdp_cuda,
flash_attention,
)
class TestSelectSdpCuda(unittest.TestCase):
"""Test _select_sdp_cuda head_dim based selection."""
def test_small_head_dim(self):
"""head_dim <= 256 should select flash_attn."""
result = _select_sdp_cuda(64)
self.assertEqual(result, "flash_attn")
def test_large_head_dim(self):
"""head_dim > 256 should select mem_efficient."""
result = _select_sdp_cuda(512)
self.assertEqual(result, "mem_efficient")
def test_boundary_head_dim(self):
"""head_dim == 256 should select flash_attn."""
result = _select_sdp_cuda(256)
self.assertEqual(result, "flash_attn")
def test_just_above_boundary(self):
"""head_dim == 257 should select mem_efficient."""
result = _select_sdp_cuda(257)
self.assertEqual(result, "mem_efficient")
class TestSelectSdp(unittest.TestCase):
"""Test _select_sdp backend selection."""
def test_select_sdp_returns_string(self):
"""_select_sdp should return a string backend name."""
try:
result = _select_sdp(64)
self.assertIsInstance(result, str)
self.assertIn(result, ["math", "flash_attn", "mem_efficient"])
except AssertionError:
# No available backend is also a valid code path
pass
class TestFlashAttentionForward(unittest.TestCase):
"""Test flash_attention forward."""
def setUp(self):
paddle.disable_static()
def test_flash_attention_basic_cpu(self):
"""Basic flash_attention on CPU (may use math backend)."""
q = paddle.randn([2, 8, 4, 16], dtype='float32')
k = paddle.randn([2, 8, 4, 16], dtype='float32')
v = paddle.randn([2, 8, 4, 16], dtype='float32')
try:
out, _ = flash_attention(q, k, v)
self.assertEqual(out.shape, [2, 8, 4, 16])
except (AssertionError, RuntimeError):
# May not be supported on CPU without proper backend
pass
def test_flash_attention_with_dropout(self):
"""Flash attention with dropout."""
q = paddle.randn([2, 8, 4, 16], dtype='float32')
k = paddle.randn([2, 8, 4, 16], dtype='float32')
v = paddle.randn([2, 8, 4, 16], dtype='float32')
try:
out, _ = flash_attention(q, k, v, dropout=0.1, training=True)
self.assertEqual(out.shape, [2, 8, 4, 16])
except (AssertionError, RuntimeError):
# May not be supported on CPU
pass
def test_flash_attention_causal(self):
"""Flash attention with causal mask."""
q = paddle.randn([2, 8, 4, 16], dtype='float32')
k = paddle.randn([2, 8, 4, 16], dtype='float32')
v = paddle.randn([2, 8, 4, 16], dtype='float32')
try:
out, _ = flash_attention(q, k, v, causal=True)
self.assertEqual(out.shape, [2, 8, 4, 16])
except (AssertionError, RuntimeError):
pass
def test_scaled_dot_product_attention_basic(self):
"""Scaled dot product attention basic.
QKV layout: [batch_size, seq_len, num_heads, head_dim]
"""
q = paddle.randn([2, 8, 4, 16], dtype='float32')
k = paddle.randn([2, 8, 4, 16], dtype='float32')
v = paddle.randn([2, 8, 4, 16], dtype='float32')
try:
out = paddle.nn.functional.scaled_dot_product_attention(q, k, v)
self.assertEqual(out.shape, [2, 8, 4, 16])
except Exception:
pass
def test_scaled_dot_product_attention_causal(self):
"""Scaled dot product attention with causal mask."""
q = paddle.randn([2, 8, 4, 16], dtype='float32')
k = paddle.randn([2, 8, 4, 16], dtype='float32')
v = paddle.randn([2, 8, 4, 16], dtype='float32')
try:
out = paddle.nn.functional.scaled_dot_product_attention(
q, k, v, is_causal=True
)
self.assertEqual(out.shape, [2, 8, 4, 16])
except Exception:
pass
if __name__ == '__main__':
unittest.main()