custom
code
sovereign-compute
snapkitty-transformer / kernels /attention_kernel.py
SNAPKITTYWEST's picture
chore: push from SNAPKITTYWEST local build
fd6abd3 verified
Raw
History Blame Contribute Delete
4.42 kB
import triton
import triton.language as tl
import torch
@triton.jit
def _flash_attn_fwd_kernel(
Q, K, V, O,
batch_size, num_heads, seq_len, head_dim,
stride_qb, stride_qh, stride_qd,
stride_kb, stride_kh, stride_kd,
stride_vb, stride_vh, stride_vd,
stride_ob, stride_oh, stride_od,
scale,
causal: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
sliding_window: tl.constexpr = -1,
num_sinks: tl.constexpr = 0,
):
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
pid_m = tl.program_id(2)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
q_ptr = Q + pid_b * stride_qb + pid_h * stride_qh
k_ptr = K + pid_b * stride_kb + pid_h * stride_kh
v_ptr = V + pid_b * stride_vb + pid_h * stride_vh
o_ptr = O + pid_b * stride_ob + pid_h * stride_oh
q = tl.load(q_ptr + offs_m[:, None] * stride_qd + offs_k[None, :],
mask=offs_m[:, None] < seq_len, other=0.0).to(tl.float32)
m_prev = tl.full([BLOCK_M], value=-1e9, dtype=tl.float32)
l_prev = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, BLOCK_K], dtype=tl.float32)
for start_n in range(0, seq_len, BLOCK_N):
offs_n_cur = start_n + offs_n
k = tl.load(k_ptr + offs_n_cur[:, None] * stride_kb + offs_k[None, :],
mask=offs_n_cur[:, None] < seq_len, other=0.0).to(tl.float32)
v = tl.load(v_ptr + offs_n_cur[:, None] * stride_vb + offs_k[None, :],
mask=offs_n_cur[:, None] < seq_len, other=0.0).to(tl.float32)
qk = tl.dot(q, tl.trans(k)) * scale
if causal:
mask = offs_m[:, None] >= offs_n_cur[None, :]
if sliding_window > 0:
mask = mask & (offs_m[:, None] - offs_n_cur[None, :] < sliding_window)
if num_sinks > 0:
sink_mask = offs_n_cur[None, :] < num_sinks
mask = mask | sink_mask
qk = tl.where(mask, qk, -1e9)
m_cur = tl.max(qk, axis=1)
m_new = tl.maximum(m_prev, m_cur)
alpha = tl.exp(m_prev - m_new)
beta = tl.exp(m_cur - m_new)
p = tl.exp(qk - m_new[:, None])
l_cur = tl.sum(p, axis=1)
l_new = alpha * l_prev + beta * l_cur
acc = acc * (alpha / l_new)[:, None] + tl.dot(p, v) * (1.0 / l_new)[:, None]
m_prev = m_new
l_prev = l_new
acc = acc.to(O.dtype.element_ty)
tl.store(o_ptr + offs_m[:, None] * stride_od + offs_k[None, :], acc,
mask=offs_m[:, None] < seq_len)
def flash_attention_triton(
Q, K, V,
causal=False,
scale=None,
sliding_window=-1,
num_sinks=0
):
batch, heads, seq_len, head_dim = Q.shape
if scale is None:
scale = head_dim ** -0.5
O = torch.zeros_like(Q)
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = head_dim
grid = (batch * heads, 1, triton.cdiv(seq_len, BLOCK_M))
_flash_attn_fwd_kernel[grid](
Q, K, V, O,
batch, heads, seq_len, head_dim,
Q.stride(0), Q.stride(1), Q.stride(3),
K.stride(0), K.stride(1), K.stride(3),
V.stride(0), V.stride(1), V.stride(3),
O.stride(0), O.stride(1), O.stride(3),
scale,
causal=causal,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
sliding_window=sliding_window,
num_sinks=num_sinks,
)
return O
def test_correctness():
torch.manual_seed(42)
for seq_len in [256, 512, 1024, 2048]:
for causal in [False, True]:
Q = torch.randn(2, 4, seq_len, 64, dtype=torch.float16, device='cuda')
K = torch.randn(2, 4, seq_len, 64, dtype=torch.float16, device='cuda')
V = torch.randn(2, 4, seq_len, 64, dtype=torch.float16, device='cuda')
O_tri = flash_attention_triton(Q, K, V, causal=causal)
O_ref = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=causal)
max_diff = (O_tri.float() - O_ref.float()).abs().max().item()
print(f"N={seq_len} causal={causal}: max_diff={max_diff:.6f}")
assert max_diff < 1e-3, f"FAILED: max_diff={max_diff}"
if __name__ == '__main__':
test_correctness()