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()