import torch
import torch.nn.functional as F
import math

def flash_attention_forward(Q, K, V, SRAM_size_M: int, scale=None):
    """
    教学版 FlashAttention forward
    Q: [N, d]
    K: [N, d]
    V: [N, dv]
    """
    N, d = Q.shape
    dv = V.shape[1]

    if scale is None:
        scale = 1.0  # 若要对齐 Transformer,可改为 1.0 / math.sqrt(d)

    block_size_base = SRAM_size_M // (4 * d)
    Bc = min(block_size_base, N)
    Br = min(block_size_base, N)
    assert Bc >= 1 and Br >= 1, "SRAM配置错误"

    O = torch.zeros((N, dv), device=Q.device, dtype=Q.dtype)
    m = torch.full((N,), -float('inf'), device=Q.device, dtype=Q.dtype)
    l = torch.zeros((N,), device=Q.device, dtype=Q.dtype)

    for j in range(0, N, Bc):
        j_end = min(j + Bc, N)
        Kj = K[j:j_end]          # [Bc, d]
        Vj = V[j:j_end]          # [Bc, dv]

        for i in range(0, N, Br):
            i_end = min(i + Br, N)
            Qi = Q[i:i_end]      # [Br, d]
            mi = m[i:i_end]      # [Br]
            li = l[i:i_end]      # [Br]
            Oi = O[i:i_end]      # [Br, dv]

            # 块内分数
            Sij = (Qi @ Kj.T) * scale              # [Br, Bc]

            # 块内 softmax 统计量
            mij = torch.max(Sij, dim=-1).values    # [Br]
            Pij = torch.exp(Sij - mij.unsqueeze(-1))
            lij = torch.sum(Pij, dim=-1)           # [Br]

            # 全局融合
            mi_new = torch.maximum(mi, mij)        # [Br]
            alpha = torch.exp(mi - mi_new)         # [Br]
            beta = torch.exp(mij - mi_new)         # [Br]
            li_new = alpha * li + beta * lij       # [Br]

            Oi_new = (
                ((alpha * li).unsqueeze(-1) * Oi) +
                (beta.unsqueeze(-1) * (Pij @ Vj))
            ) / li_new.unsqueeze(-1)

            O[i:i_end] = Oi_new
            m[i:i_end] = mi_new
            l[i:i_end] = li_new
            
            print('O=',O)
            
            print('*'*30+'\n')

    return O, m, l


def standard_attention(Q, K, V, scale=None):
    d = Q.shape[1]
    if scale is None:
        scale = 1.0  # 若要对齐 Transformer,可改为 1.0 / math.sqrt(d)
    S = (Q @ K.T) * scale
    P = F.softmax(S, dim=-1)
    return P @ V


if __name__ == '__main__':
    N = 1024
    d = 64
    SRAM_size_M = 1024 * 16

    torch.manual_seed(42)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    Q = torch.randn(N, d, device=device)
    K = torch.randn(N, d, device=device)
    V = torch.randn(N, d, device=device)

    O_flash, m, l = flash_attention_forward(Q, K, V, SRAM_size_M)
    O_std = standard_attention(Q, K, V)

    diff = torch.max(torch.abs(O_flash - O_std))
    print(f"max error = {diff.item():.8e}")
    print("一致" if diff < 1e-5 else "不一致")
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐