Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions benchmarks/operators/block_sparse_attention/config.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
test_cases:
- B: 2
H: 8
M: 1024
D: 64
H_kv: 2
BLOCK_M: 64
BLOCK_N: 64
BLOCK_D: 64
NUM_D_BLOCKS: 1
dtype: "float16"

- B: 1
H: 16
M: 4096
D: 128
H_kv: 4
BLOCK_M: 128
BLOCK_N: 128
BLOCK_D: 64
NUM_D_BLOCKS: 2
dtype: "float16"
137 changes: 137 additions & 0 deletions benchmarks/operators/block_sparse_attention/impl_cutile.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
import torch
import cuda.tile as ct
import math

ConstInt = ct.Constant[int]

@ct.kernel
def block_sparse_attention_cutile_kernel(
Out, Q, K, V,
csr_row_indices, csr_col_indices,
csr_row_stride_h: ConstInt, csr_col_stride_h: ConstInt,
num_layout: ConstInt, softmax_scale: ct.Constant[float],
num_heads: ConstInt, num_kv_heads: ConstInt, total_seq_len: ConstInt,
BLOCK_M: ConstInt, BLOCK_N: ConstInt, BLOCK_D: ConstInt
):
start_m = ct.bid(0)
off_bh = ct.bid(1)

off_h = off_bh % num_heads
off_b = off_bh // num_heads

# GQA mapping
head_groups = num_heads // num_kv_heads
off_h_kv = off_h // head_groups

# Load Q tile and reshape to 2D
q_tile = ct.load(Q, index=(off_b, off_h, start_m, 0), shape=(1, 1, BLOCK_M, BLOCK_D), padding_mode=ct.PaddingMode.ZERO)
q_tile_2d = ct.reshape(q_tile, (BLOCK_M, BLOCK_D))

# Initialize Online Softmax states
m_i = ct.full((BLOCK_M, 1), -float('inf'), dtype=ct.float32)
l_i = ct.full((BLOCK_M, 1), 0.0, dtype=ct.float32)
acc = ct.full((BLOCK_M, BLOCK_D), 0.0, dtype=ct.float32)

# Fetch CSR pointers
layout_h = off_h % num_layout
row_idx_ptr = layout_h * csr_row_stride_h + start_m
start_l = ct.load(csr_row_indices, index=(row_idx_ptr,), shape=())
end_l = ct.load(csr_row_indices, index=(row_idx_ptr + 1,), shape=())

# Setup row offsets for masking
offs_m = start_m * BLOCK_M + ct.expand_dims(ct.arange(BLOCK_M, dtype=ct.int32), 1)
valid_m = offs_m < total_seq_len

# Sparse Loop
l = start_l
while l < end_l:
col_idx_ptr = layout_h * csr_col_stride_h + l
col_idx = ct.load(csr_col_indices, index=(col_idx_ptr,), shape=())
start_n = col_idx

# Load K and compute QK
k_tile = ct.load(K, index=(off_b, off_h_kv, start_n, 0), shape=(1, 1, BLOCK_N, BLOCK_D), padding_mode=ct.PaddingMode.ZERO)
k_tile_2d = ct.reshape(k_tile, (BLOCK_N, BLOCK_D))
k_tile_T = ct.transpose(k_tile_2d, 0, 1)

qk = ct.full((BLOCK_M, BLOCK_N), 0.0, dtype=ct.float32)
qk = ct.mma(q_tile_2d, k_tile_T, qk)
qk = qk * softmax_scale

# Causal & Sequence Masking
offs_n = start_n * BLOCK_N + ct.expand_dims(ct.arange(BLOCK_N, dtype=ct.int32), 0)
valid_n = offs_n < total_seq_len
causal_mask = offs_m >= offs_n
seq_mask = ct.bitwise_and(valid_m, valid_n)
mask = ct.bitwise_and(causal_mask, seq_mask)

qk = ct.where(mask, qk, -float('inf'))

# Online Softmax Math
m_ij = ct.max(qk, axis=1)
m_ij = ct.expand_dims(m_ij, 1)

m_i_new = ct.maximum(m_i, m_ij)
alpha = ct.exp(m_i - m_i_new)
beta = ct.exp(m_ij - m_i_new)

p = ct.exp(qk - m_ij)
l_ij = ct.sum(p, axis=1)
l_ij = ct.expand_dims(l_ij, 1)

l_i_new = alpha * l_i + beta * l_ij

p_scale = beta / l_i_new
p = p * p_scale

acc_scale = (l_i / l_i_new) * alpha
acc = acc * acc_scale

# Downcast P and multiply with V
p_casted = ct.astype(p, Q.dtype)

v_tile = ct.load(V, index=(off_b, off_h_kv, start_n, 0), shape=(1, 1, BLOCK_N, BLOCK_D), padding_mode=ct.PaddingMode.ZERO)
v_tile_2d = ct.reshape(v_tile, (BLOCK_N, BLOCK_D))

acc = ct.mma(p_casted, v_tile_2d, acc)

# Update states
l_i = l_i_new
m_i = m_i_new

l = l + 1

# Write back result
acc_casted = ct.astype(acc, Out.dtype)
acc_reshaped = ct.reshape(acc_casted, (1, 1, BLOCK_M, BLOCK_D))
ct.store(Out, index=(off_b, off_h, start_m, 0), tile=acc_reshaped)

def run(Q, K, V, layout_csr_row_indices, layout_csr_col_indices,
layout_csr_row_stride_h, layout_csr_col_stride_h,
num_layout, softmax_scale, num_heads, num_kv_heads,
total_seq_len, BLOCK_M, EVEN_M, BLOCK_N, EVEN_N, BLOCK_D, NUM_D_BLOCKS,
block_size: int = None):

batch_size = Q.shape[0]
out = torch.empty_like(Q)

# Grid:[M_blocks, Batch * Heads, 1]
grid = (math.ceil(total_seq_len / BLOCK_M), batch_size * num_heads, 1)

# cuTile handles D dimension directly.
# If Triton used NUM_D_BLOCKS=2 (e.g. D=128, BLOCK_D=64), we can just tell cuTile to load 128.
ACTUAL_BLOCK_D = BLOCK_D * NUM_D_BLOCKS

ct.launch(
torch.cuda.current_stream(),
grid,
block_sparse_attention_cutile_kernel,
(out, Q, K, V,
layout_csr_row_indices, layout_csr_col_indices,
layout_csr_row_stride_h, layout_csr_col_stride_h,
num_layout, float(softmax_scale),
num_heads, num_kv_heads, total_seq_len,
BLOCK_M, BLOCK_N, ACTUAL_BLOCK_D)
)

return out
59 changes: 59 additions & 0 deletions benchmarks/operators/block_sparse_attention/impl_torch.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
import torch
import math

def run(Q, K, V, layout_csr_row_indices, layout_csr_col_indices,
layout_csr_row_stride_h, layout_csr_col_stride_h,
num_layout, softmax_scale, num_heads, num_kv_heads,
total_seq_len, BLOCK_M, EVEN_M, BLOCK_N, EVEN_N, BLOCK_D, NUM_D_BLOCKS):
"""
Reference PyTorch implementation for Block Sparse Attention.
"""
B, H, M, D = Q.shape
_, H_kv, N, _ = K.shape

# 1. Handle GQA (Grouped Query Attention) by repeating K and V
head_groups = H // H_kv
#[B, H_kv, N, D] -> [B, H_kv, head_groups, N, D] ->[B, H, N, D]
K_expanded = K.unsqueeze(2).expand(B, H_kv, head_groups, N, D).reshape(B, H, N, D)
V_expanded = V.unsqueeze(2).expand(B, H_kv, head_groups, N, D).reshape(B, H, N, D)

# 2. Reconstruct the dense attention mask from CSR representation
num_rows = math.ceil(M / BLOCK_M)
num_cols = math.ceil(N / BLOCK_N)

# Initialize with -inf (masked out)
sparse_mask = torch.full((H, M, N), float('-inf'), device=Q.device, dtype=torch.float32)

for h in range(H):
layout_h = h % num_layout
for r in range(num_rows):
start_l = layout_csr_row_indices[layout_h * layout_csr_row_stride_h + r].item()
end_l = layout_csr_row_indices[layout_h * layout_csr_row_stride_h + r + 1].item()

for l in range(start_l, end_l):
c = layout_csr_col_indices[layout_h * layout_csr_col_stride_h + l].item()
r_start, r_end = r * BLOCK_M, min((r + 1) * BLOCK_M, M)
c_start, c_end = c * BLOCK_N, min((c + 1) * BLOCK_N, N)

# Unmask this block
sparse_mask[h, r_start:r_end, c_start:c_end] = 0.0

# Expand mask for batch size
sparse_mask = sparse_mask.unsqueeze(0).expand(B, H, M, N)

# 3. Create Causal Mask (Lower Triangular)
# The kernel says: qk += tl.where(offs_m[:, None] >= (start_n + offs_n[None, :]), 0, float("-inf"))
causal_mask = torch.tril(torch.ones(M, N, device=Q.device, dtype=torch.bool))
causal_mask = torch.where(causal_mask, 0.0, float('-inf'))

# 4. Dense Attention Computation
scores = torch.matmul(Q.float(), K_expanded.float().transpose(-2, -1)) * softmax_scale

# Apply both Sparse Mask and Causal Mask
scores = scores + sparse_mask + causal_mask.unsqueeze(0).unsqueeze(0)

probs = torch.softmax(scores, dim=-1)

out = torch.matmul(probs, V_expanded.float())

return out.to(Q.dtype)
Loading