Download kernels/attention_kernel.py from Snapkitty/snapkitty-transformer: direct link, hf CLI and curl.
- Browser
- Download file 4.42 kB
-
https://huggingface.co/Snapkitty/snapkitty-transformer/resolve/main/kernels/attention_kernel.py
- Command line
-
hf download hf://Snapkitty/snapkitty-transformer/kernels/attention_kernel.py
-
curl -L -o attention_kernel.py https://huggingface.co/Snapkitty/snapkitty-transformer/resolve/main/kernels/attention_kernel.py
4.42 kB
| import triton | |
| import triton.language as tl | |
| import torch | |
| 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() |