replicate/flash-attn2
0165
1# /// script2# dependencies = [3# "numpy", 4# "torch", 5# "kernels"6# ]7# ///8import torch9from kernels import get_kernel10 11# Setup12torch.manual_seed(42)13flash_attn = get_kernel("kernels-community/flash-attn")14device = torch.device("cuda")15 16# Create test tensors17B, S, H, D = 2, 5, 4, 8 # batch, seq_len, heads, head_dim18q = k = v = torch.randn(B, S, H, D, device=device, dtype=torch.float16)19 20# Reference implementation using PyTorch SDPA21def reference_attention(query, key, value, causal=False):22 query, key, value = (x.transpose(1, 2).contiguous() for x in (query, key, value))23 with torch.nn.attention.sdpa_kernel(torch.nn.attention.SDPBackend.MATH):24 out = torch.nn.functional.scaled_dot_product_attention(query, key, value, is_causal=causal)25 return out.transpose(1, 2).contiguous()26 27# 1. Standard attention28print("\n1. Standard attention:")29out_ref = reference_attention(q, k, v)30out_flash = flash_attn.fwd(31 q=q, 32 k=k, 33 v=v, 34 is_causal=False,35)[0]36print(f"Reference output: {out_ref.shape}")37print(f"Flash output: {out_flash.shape}")38print(f"Outputs close: {torch.allclose(out_flash, out_ref, atol=1e-2, rtol=1e-3)}")39 40# 2. Causal attention (for autoregressive models)41print("\n2. Causal attention:")42 43out_ref_causal = reference_attention(q, k, v, causal=True)44out_causal = flash_attn.fwd(45 q=q, 46 k=k, 47 v=v, 48 is_causal=True,49)[0]50print(f"Reference causal output: {out_ref_causal.shape}")51print(f"Flash causal output: {out_causal.shape}")52print(f"Outputs close: {torch.allclose(out_causal, out_ref_causal, atol=1e-2, rtol=1e-3)}")53 54def var_reference_attention(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, causal=False):55 batch_size = cu_seqlens_q.shape[0] - 156 # Return output in packed format (same as flash attention)57 total_tokens_q = q.shape[0]58 out = torch.zeros((total_tokens_q, q.shape[1], q.shape[2]), device=q.device, dtype=q.dtype)59 60 for b in range(batch_size):61 start_q, end_q = cu_seqlens_q[b], cu_seqlens_q[b + 1]62 start_k, end_k = cu_seqlens_k[b], cu_seqlens_k[b + 1]63 64 # Extract slices for this batch65 q_slice = q[start_q:end_q] # Shape: (seq_len_q, H, D)66 k_slice = k[start_k:end_k] # Shape: (seq_len_k, H, D)67 v_slice = v[start_k:end_k] # Shape: (seq_len_k, H, D)68 69 # Add batch dimension for reference_attention70 q_slice = q_slice.unsqueeze(0) # Shape: (1, seq_len_q, H, D)71 k_slice = k_slice.unsqueeze(0) # Shape: (1, seq_len_k, H, D)72 v_slice = v_slice.unsqueeze(0) # Shape: (1, seq_len_k, H, D)73 74 # Compute attention and remove batch dimension75 attn_out = reference_attention(q_slice, k_slice, v_slice, causal=causal)76 attn_out = attn_out.squeeze(0) # Shape: (seq_len_q, H, D)77 78 # Place result in output tensor (packed format)79 out[start_q:end_q] = attn_out80 81 return out82 83# 3. Variable length sequences (packed format)84print("\n3. Variable length sequences:")85# Pack sequences of lengths [3,4,3] for q and [4,5,3] for k into single tensors86q_var = torch.randn(10, H, D, device=device, dtype=torch.float16) # total_q=1087k_var = v_var = torch.randn(12, H, D, device=device, dtype=torch.float16) # total_k=1288cu_q = torch.tensor([0, 3, 7, 10], device=device, dtype=torch.int32) # cumulative sequence lengths89cu_k = torch.tensor([0, 4, 9, 12], device=device, dtype=torch.int32)90 91out_var_ref = var_reference_attention(q_var, k_var, v_var, cu_q, cu_k, max_seqlen_q=4, max_seqlen_k=5, causal=False)92# Custom function to handle variable93out_var = flash_attn.varlen_fwd(94 q=q_var,95 k=k_var,96 v=v_var,97 cu_seqlens_q=cu_q,98 cu_seqlens_k=cu_k,99 max_seqlen_q=4,100 max_seqlen_k=5,101)[0]102print(f"Variable length output: {out_var.shape}")103print(f"Reference variable length output: {out_var_ref.shape}")104print(f"Outputs close: {torch.allclose(out_var, out_var_ref, atol=1e-2, rtol=1e-3)}")105 