Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e9f0f714ed | ||
|
|
554a8f045e | ||
|
|
6d0eb1b956 | ||
|
|
20f5d1a779 | ||
|
|
5da1c9aadf |
+1
-1
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
path = sta_kernel/thunderkitten/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -0,0 +1,252 @@
|
||||
# Copyright (c) Tile-AI Corporation.
|
||||
# Licensed under the MIT License.
|
||||
import math
|
||||
import torch
|
||||
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
import torch.nn.functional as F
|
||||
|
||||
def get_sta_mask(x, canvas_size=(32, 48, 80), tile_size=(4, 8, 8), kernel_size=(1, 1, 1), has_text=False):
|
||||
bsz, num_head, downsample_len, _ = x.shape
|
||||
device = x.device
|
||||
CT, CH, CW = canvas_size
|
||||
TT, TH, TW = tile_size
|
||||
NT, NH, NW = CT // TT, CH // TH, CW // TW
|
||||
KT, KH, KW = kernel_size
|
||||
DT, DH, DW = KT // 2, KH // 2, KW // 2
|
||||
|
||||
dense_mask = torch.full([bsz, num_head, downsample_len, downsample_len],
|
||||
False,
|
||||
dtype=torch.bool,
|
||||
device=device)
|
||||
indices = torch.arange(downsample_len, device=device)
|
||||
q_t = (indices // (NT * NH * NW)).unsqueeze(-1)
|
||||
q_h = ((indices // (NW)) % NH).unsqueeze(-1)
|
||||
q_w = ((indices // 1) % NW).unsqueeze(-1)
|
||||
|
||||
q_t = torch.clamp(q_t, DT, NT-DT-1)
|
||||
q_h = torch.clamp(q_h, DH, NH-DH-1)
|
||||
q_w = torch.clamp(q_w, DW, NW-DW-1)
|
||||
|
||||
k_indices = torch.arange(downsample_len, device=device)
|
||||
k_t = (k_indices // (NT * NH * NW)).unsqueeze(0)
|
||||
k_h = ((k_indices // (NW)) % NH).unsqueeze(0)
|
||||
k_w = ((k_indices // 1) % NW).unsqueeze(0)
|
||||
|
||||
t_dist = torch.abs(q_t - k_t)
|
||||
h_dist = torch.abs(q_h - k_h)
|
||||
w_dist = torch.abs(q_w - k_w)
|
||||
mask = (t_dist <= DT) & (h_dist <= DH) & (w_dist <= DW)
|
||||
|
||||
for b in range(bsz):
|
||||
for h in range(num_head):
|
||||
dense_mask[b, h] = mask
|
||||
|
||||
# for text mask
|
||||
if has_text:
|
||||
text_start = downsample_len - 3
|
||||
dense_mask[:, :, text_start:, :text_start] = True
|
||||
dense_mask[:, :, text_start:, text_start:] = torch.tril(
|
||||
torch.ones(3, 3, device=device, dtype=torch.bool)
|
||||
)
|
||||
return dense_mask
|
||||
|
||||
def blocksparse_flashattn(batch, heads, seq_len, dim, downsample_len, is_causal):
|
||||
block_M = 64
|
||||
block_N = 64
|
||||
num_stages = 1
|
||||
threads = 128
|
||||
scale = (1.0 / dim)**0.5 * 1.44269504 # log2(e)
|
||||
shape = [batch, heads, seq_len, dim]
|
||||
block_mask_shape = [batch, heads, downsample_len, downsample_len]
|
||||
|
||||
dtype = "float16"
|
||||
accum_dtype = "float"
|
||||
block_mask_dtype = "bool"
|
||||
|
||||
def kernel_func(block_M, block_N, num_stages, threads):
|
||||
|
||||
@T.macro
|
||||
def MMA0(
|
||||
K: T.Buffer(shape, dtype),
|
||||
Q_shared: T.Buffer([block_M, dim], dtype),
|
||||
K_shared: T.Buffer([block_N, dim], dtype),
|
||||
acc_s: T.Buffer([block_M, block_N], accum_dtype),
|
||||
k: T.int32,
|
||||
bx: T.int32,
|
||||
by: T.int32,
|
||||
bz: T.int32,
|
||||
):
|
||||
T.copy(K[bz, by, k * block_N:(k + 1) * block_N, :], K_shared)
|
||||
if is_causal:
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
acc_s[i, j] = T.if_then_else(bx * block_M + i >= k * block_N + j, 0,
|
||||
-T.infinity(acc_s.dtype))
|
||||
else:
|
||||
T.clear(acc_s)
|
||||
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullRow)
|
||||
|
||||
@T.macro
|
||||
def MMA1(
|
||||
V: T.Buffer(shape, dtype),
|
||||
V_shared: T.Buffer([block_M, dim], dtype),
|
||||
acc_s_cast: T.Buffer([block_M, block_N], dtype),
|
||||
acc_o: T.Buffer([block_M, dim], accum_dtype),
|
||||
k: T.int32,
|
||||
by: T.int32,
|
||||
bz: T.int32,
|
||||
):
|
||||
T.copy(V[bz, by, k * block_N:(k + 1) * block_N, :], V_shared)
|
||||
T.gemm(acc_s_cast, V_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)
|
||||
|
||||
@T.macro
|
||||
def Softmax(
|
||||
acc_s: T.Buffer([block_M, block_N], accum_dtype),
|
||||
acc_s_cast: T.Buffer([block_M, block_N], dtype),
|
||||
scores_max: T.Buffer([block_M], accum_dtype),
|
||||
scores_max_prev: T.Buffer([block_M], accum_dtype),
|
||||
scores_scale: T.Buffer([block_M], accum_dtype),
|
||||
scores_sum: T.Buffer([block_M], accum_dtype),
|
||||
logsum: T.Buffer([block_M], accum_dtype),
|
||||
):
|
||||
T.copy(scores_max, scores_max_prev)
|
||||
T.fill(scores_max, -T.infinity(accum_dtype))
|
||||
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
|
||||
# To do causal softmax, we need to set the scores_max to 0 if it is -inf
|
||||
# This process is called Check_inf in FlashAttention3 code, and it only need to be done
|
||||
# in the first ceil_div(kBlockM, kBlockN) steps.
|
||||
# for i in T.Parallel(block_M):
|
||||
# scores_max[i] = T.if_then_else(scores_max[i] == -T.infinity(accum_dtype), 0, scores_max[i])
|
||||
for i in T.Parallel(block_M):
|
||||
scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale)
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
# Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
|
||||
# max * log_2(e)) This allows the compiler to use the ffma
|
||||
# instruction instead of fadd and fmul separately.
|
||||
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
|
||||
T.reduce_sum(acc_s, scores_sum, dim=1)
|
||||
for i in T.Parallel(block_M):
|
||||
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
|
||||
T.copy(acc_s, acc_s_cast)
|
||||
|
||||
@T.macro
|
||||
def Rescale(
|
||||
acc_o: T.Buffer([block_M, dim], accum_dtype),
|
||||
scores_scale: T.Buffer([block_M], accum_dtype),
|
||||
):
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] *= scores_scale[i]
|
||||
|
||||
@T.prim_func
|
||||
def main(
|
||||
Q: T.Buffer(shape, dtype),
|
||||
K: T.Buffer(shape, dtype),
|
||||
V: T.Buffer(shape, dtype),
|
||||
BlockSparseMask: T.Buffer(block_mask_shape, block_mask_dtype),
|
||||
Output: T.Buffer(shape, dtype),
|
||||
):
|
||||
with T.Kernel(
|
||||
T.ceildiv(seq_len, block_M), heads, batch, threads=threads) as (bx, by, bz):
|
||||
Q_shared = T.alloc_shared([block_M, dim], dtype)
|
||||
K_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
V_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
O_shared = T.alloc_shared([block_M, dim], dtype)
|
||||
acc_s = T.alloc_fragment([block_M, block_N], accum_dtype)
|
||||
acc_s_cast = T.alloc_fragment([block_M, block_N], dtype)
|
||||
acc_o = T.alloc_fragment([block_M, dim], accum_dtype)
|
||||
scores_max = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_max_prev = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_scale = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_sum = T.alloc_fragment([block_M], accum_dtype)
|
||||
logsum = T.alloc_fragment([block_M], accum_dtype)
|
||||
block_mask = T.alloc_local([downsample_len], block_mask_dtype)
|
||||
|
||||
T.copy(Q[bz, by, bx * block_M:(bx + 1) * block_M, :], Q_shared)
|
||||
T.fill(acc_o, 0)
|
||||
T.fill(logsum, 0)
|
||||
T.fill(scores_max, -T.infinity(accum_dtype))
|
||||
|
||||
for vj in T.serial(downsample_len):
|
||||
block_mask[vj] = BlockSparseMask[bz, by, bx, vj]
|
||||
|
||||
loop_range = (
|
||||
T.min(T.ceildiv(seq_len, block_N), T.ceildiv(
|
||||
(bx + 1) * block_M, block_N)) if is_causal else T.ceildiv(seq_len, block_N))
|
||||
|
||||
for k in T.Pipelined(loop_range, num_stages=num_stages):
|
||||
if block_mask[k] != 0:
|
||||
MMA0(K, Q_shared, K_shared, acc_s, k, bx, by, bz)
|
||||
Softmax(acc_s, acc_s_cast, scores_max, scores_max_prev, scores_scale,
|
||||
scores_sum, logsum)
|
||||
Rescale(acc_o, scores_scale)
|
||||
MMA1(V, V_shared, acc_s_cast, acc_o, k, by, bz)
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] /= logsum[i]
|
||||
T.copy(acc_o, O_shared)
|
||||
T.copy(O_shared, Output[bz, by, bx * block_M:(bx + 1) * block_M, :])
|
||||
|
||||
return main
|
||||
|
||||
return kernel_func(block_M, block_N, num_stages, threads)
|
||||
|
||||
|
||||
def test_sta_attention():
|
||||
# Config
|
||||
BATCH, N_HEADS, SEQ_LEN, D_HEAD = 1, 24, 2048, 128
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Create inputs
|
||||
q = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
k = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
v = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
|
||||
sm_scale = 1.0 / (D_HEAD**0.5)
|
||||
|
||||
# Create sparse mask (downsampled to block level)
|
||||
canvas_size = (32, 48, 80)
|
||||
tile_size = (4, 8, 8)
|
||||
delta_size = (2, 2, 2)
|
||||
BLOCK = tile_size[0] * tile_size[1] * tile_size[2]
|
||||
downsample_factor = BLOCK
|
||||
downsample_len = math.ceil(SEQ_LEN / downsample_factor)
|
||||
x_ds = torch.randn([BATCH, N_HEADS, downsample_len, downsample_len],
|
||||
device='cuda',
|
||||
dtype=torch.bfloat16)
|
||||
block_mask = get_sta_mask(x_ds, canvas_size, tile_size, delta_size, has_text=True)
|
||||
# print mask density
|
||||
print("mask density", block_mask.sum() / block_mask.numel())
|
||||
print("block_mask", block_mask)
|
||||
|
||||
# Run Triton kernel
|
||||
program = blocksparse_flashattn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, downsample_len, is_causal=True)
|
||||
kernel = tilelang.compile(program, out_idx=[4])
|
||||
|
||||
cuda_source = kernel.get_kernel_source()
|
||||
print("Generated CUDA kernel:\n", cuda_source)
|
||||
|
||||
tilelang_output = kernel(q, k, v, block_mask)
|
||||
|
||||
if True:
|
||||
# Compute reference
|
||||
# Expand block mask to full attention matrix
|
||||
full_mask = torch.kron(block_mask.float(), torch.ones(BLOCK, BLOCK, device='cuda'))
|
||||
full_mask = full_mask[..., :SEQ_LEN, :SEQ_LEN].bool()
|
||||
full_mask = full_mask & torch.tril(torch.ones_like(full_mask)) # Apply causal
|
||||
|
||||
# PyTorch reference implementation
|
||||
attn = torch.einsum('bhsd,bhtd->bhst', q, k) * sm_scale
|
||||
attn = attn.masked_fill(~full_mask, float('-inf'))
|
||||
attn = F.softmax(attn, dim=-1)
|
||||
ref_output = torch.einsum('bhst,bhtd->bhsd', attn, v)
|
||||
|
||||
print("ref_output", ref_output)
|
||||
print("tilelang_output", tilelang_output)
|
||||
|
||||
# Verify accuracy
|
||||
torch.testing.assert_close(tilelang_output, ref_output, atol=1e-2, rtol=1e-2)
|
||||
print("Pass topk sparse attention test with qlen == klen")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_sta_attention()
|
||||
Reference in New Issue
Block a user