From 9083918640d0a7a0f43250dcbd86309b3569b530 Mon Sep 17 00:00:00 2001 From: zzczzc20 Date: Mon, 16 Mar 2026 10:20:44 +0000 Subject: [PATCH] Reapply denton_dev code onto latest main --- .../block_sparse_attention/config.yaml | 22 ++ .../block_sparse_attention/impl_cutile.py | 137 +++++++++++ .../block_sparse_attention/impl_torch.py | 59 +++++ .../block_sparse_attention/impl_triton.py | 192 ++++++++++++++++ .../operators/flash_attention/config.yaml | 16 ++ .../operators/flash_attention/impl_cutile.py | 158 +++++++++++++ .../operators/flash_attention/impl_torch.py | 5 + .../operators/flash_attention/impl_triton.py | 122 ++++++++++ benchmarks/operators/flash_decode/config.yaml | 24 ++ .../operators/flash_decode/impl_cutile.py | 85 +++++++ .../operators/flash_decode/impl_torch.py | 78 +++++++ .../operators/flash_decode/impl_triton.py | 109 +++++++++ .../operators/matmul_fp16_fp8/config.yaml | 9 + .../operators/matmul_fp16_fp8/impl_cutile.py | 77 +++++++ .../operators/matmul_fp16_fp8/impl_torch.py | 19 ++ .../operators/matmul_fp16_fp8/impl_triton.py | 128 +++++++++++ benchmarks/operators/matmul_int8/config.yaml | 20 ++ .../operators/matmul_int8/impl_cutile.py | 75 ++++++ .../operators/matmul_int8/impl_torch.py | 21 ++ .../operators/matmul_int8/impl_triton.py | 105 +++++++++ benchmarks/operators/rope/config.yaml | 37 +++ benchmarks/operators/rope/impl_cutile.py | 56 +++++ benchmarks/operators/rope/impl_torch.py | 23 ++ benchmarks/operators/rope/impl_triton.py | 94 ++++++++ benchmarks/operators/softmax/config.yaml | 20 ++ benchmarks/operators/softmax/impl_cutile.py | 53 +++++ benchmarks/operators/softmax/impl_torch.py | 4 + benchmarks/operators/softmax/impl_triton.py | 53 +++++ .../operators/streamk_matmul/config.yaml | 29 +++ .../operators/streamk_matmul/impl_cutile.py | 217 ++++++++++++++++++ .../operators/streamk_matmul/impl_torch.py | 14 ++ .../operators/streamk_matmul/impl_triton.py | 206 +++++++++++++++++ core/verifier.py | 2 +- data/tensors.py | 128 +++++++++++ 34 files changed, 2396 insertions(+), 1 deletion(-) create mode 100644 benchmarks/operators/block_sparse_attention/config.yaml create mode 100644 benchmarks/operators/block_sparse_attention/impl_cutile.py create mode 100644 benchmarks/operators/block_sparse_attention/impl_torch.py create mode 100644 benchmarks/operators/block_sparse_attention/impl_triton.py create mode 100644 benchmarks/operators/flash_attention/config.yaml create mode 100644 benchmarks/operators/flash_attention/impl_cutile.py create mode 100644 benchmarks/operators/flash_attention/impl_torch.py create mode 100644 benchmarks/operators/flash_attention/impl_triton.py create mode 100644 benchmarks/operators/flash_decode/config.yaml create mode 100644 benchmarks/operators/flash_decode/impl_cutile.py create mode 100644 benchmarks/operators/flash_decode/impl_torch.py create mode 100644 benchmarks/operators/flash_decode/impl_triton.py create mode 100644 benchmarks/operators/matmul_fp16_fp8/config.yaml create mode 100644 benchmarks/operators/matmul_fp16_fp8/impl_cutile.py create mode 100644 benchmarks/operators/matmul_fp16_fp8/impl_torch.py create mode 100644 benchmarks/operators/matmul_fp16_fp8/impl_triton.py create mode 100644 benchmarks/operators/matmul_int8/config.yaml create mode 100644 benchmarks/operators/matmul_int8/impl_cutile.py create mode 100644 benchmarks/operators/matmul_int8/impl_torch.py create mode 100644 benchmarks/operators/matmul_int8/impl_triton.py create mode 100644 benchmarks/operators/rope/config.yaml create mode 100644 benchmarks/operators/rope/impl_cutile.py create mode 100644 benchmarks/operators/rope/impl_torch.py create mode 100644 benchmarks/operators/rope/impl_triton.py create mode 100644 benchmarks/operators/softmax/config.yaml create mode 100644 benchmarks/operators/softmax/impl_cutile.py create mode 100644 benchmarks/operators/softmax/impl_torch.py create mode 100644 benchmarks/operators/softmax/impl_triton.py create mode 100644 benchmarks/operators/streamk_matmul/config.yaml create mode 100644 benchmarks/operators/streamk_matmul/impl_cutile.py create mode 100644 benchmarks/operators/streamk_matmul/impl_torch.py create mode 100644 benchmarks/operators/streamk_matmul/impl_triton.py diff --git a/benchmarks/operators/block_sparse_attention/config.yaml b/benchmarks/operators/block_sparse_attention/config.yaml new file mode 100644 index 00000000..7d0e49d5 --- /dev/null +++ b/benchmarks/operators/block_sparse_attention/config.yaml @@ -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" \ No newline at end of file diff --git a/benchmarks/operators/block_sparse_attention/impl_cutile.py b/benchmarks/operators/block_sparse_attention/impl_cutile.py new file mode 100644 index 00000000..954fbae0 --- /dev/null +++ b/benchmarks/operators/block_sparse_attention/impl_cutile.py @@ -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 \ No newline at end of file diff --git a/benchmarks/operators/block_sparse_attention/impl_torch.py b/benchmarks/operators/block_sparse_attention/impl_torch.py new file mode 100644 index 00000000..99d25ffe --- /dev/null +++ b/benchmarks/operators/block_sparse_attention/impl_torch.py @@ -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) \ No newline at end of file diff --git a/benchmarks/operators/block_sparse_attention/impl_triton.py b/benchmarks/operators/block_sparse_attention/impl_triton.py new file mode 100644 index 00000000..8a8757c5 --- /dev/null +++ b/benchmarks/operators/block_sparse_attention/impl_triton.py @@ -0,0 +1,192 @@ +import torch +import triton +import triton.language as tl + +@triton.jit +def block_sparse_attention_kernel( + out, # output [B, H, M, D]. Note that B is batch_size, H is num_heads, M is q_seq_len, and D is head_size + Q, # query [B, H, M, D] + K, # key [B, H_kv, N, D]. Note that N is max_seq_len for kv cache, H_kv is num_kv_heads + V, # value [B, H_kv, N, D] + layout_csr_row_indices, # block mask CSR format. Shape is [L, num_rows + 1] where num_rows = max_seq_len / BLOCK_M + layout_csr_col_indices, # block mask CSR format. Shape is [L, num_rows * num_cols] where num_cols = max_seq_len / BLOCK_N + layout_csr_row_stride_h, # stride per head for csr_row_indices, i.e. num_rows + 1 + layout_csr_col_stride_h, # stride per head for csr_col_indices, i.e. num_rows * num_cols + num_layout, # number of sparse layout (L) + softmax_scale, + stride_qb, + stride_qh, + stride_qm, + stride_kb, + stride_kh, + stride_kn, + stride_vb, + stride_vh, + stride_vn, + stride_ob, + stride_oh, + stride_om, + num_heads, + num_kv_heads, + total_seq_len, # Total sequence length including past sequence length and query sequence length. + BLOCK_M: tl.constexpr, # block size for q_seq_len + EVEN_M: tl.constexpr, # whether q_seq_len % BLOCK_M == 0 + BLOCK_N: tl.constexpr, # block size for k_seq_len + EVEN_N: tl.constexpr, # whether k_seq_len % BLOCK_N == 0 + BLOCK_D: tl.constexpr, # block size for D + NUM_D_BLOCKS: tl.constexpr, # number of data blocks = D / BLOCK_D +): + tl.static_print(f"{BLOCK_M=} {BLOCK_N=} {BLOCK_D=} {EVEN_M=} {EVEN_N=} {NUM_D_BLOCKS=}") + + # Past sequence length is 0 since this kernel is for prompt only. + q_seq_len = total_seq_len + + # Grid is [CDiv(q_seq_len, BLOCK_M), batch_size * num_heads] + start_m = tl.program_id(0) + off_bh = tl.program_id(1) + + off_h = off_bh % num_heads + off_b = off_bh // num_heads + + # For group query attention, map the query head index to the corresponding one for key and value. + head_groups = num_heads // num_kv_heads + off_h_kv = off_h // head_groups + + Q += off_b * stride_qb + off_h * stride_qh + K += off_b * stride_kb + off_h_kv * stride_kh + V += off_b * stride_vb + off_h_kv * stride_vh + + # Initialize offsets + offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_n = tl.arange(0, BLOCK_N) + offs_d = tl.arange(0, BLOCK_D) + off_q = offs_m[:, None] * stride_qm + offs_d[None, :] # [BLOCK_M, BLOCK_D] + off_k = offs_n[None, :] * stride_kn + offs_d[:, None] # [BLOCK_D, BLOCK_N] + off_v = offs_n[:, None] * stride_vn + offs_d[None, :] # [BLOCK_N, BLOCK_D] + + # Initialize pointers to query, key, value + q_ptrs = Q + off_q + k_ptrs = K + off_k + v_ptrs = V + off_v + + # Initialize pointer to m and l + m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf") + l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + acc = tl.zeros([BLOCK_M, BLOCK_D], dtype=tl.float32) + if NUM_D_BLOCKS >= 2: + acc2 = tl.zeros([BLOCK_M, BLOCK_D], dtype=tl.float32) + + # Load q: it will stay in SRAM throughout + if EVEN_M: + q = tl.load(q_ptrs) + if NUM_D_BLOCKS >= 2: + q2 = tl.load(q_ptrs + BLOCK_D) + else: + q = tl.load(q_ptrs, mask=offs_m[:, None] < q_seq_len) + if NUM_D_BLOCKS >= 2: + q2 = tl.load(q_ptrs + BLOCK_D, mask=offs_m[:, None] < q_seq_len) + + layout_h = off_h % num_layout + + # This assumes that past sequence length is 0, otherwise need + (past_seq_len + 1) // BLOCK_M. + layout_ptr = layout_csr_row_indices + layout_h * layout_csr_row_stride_h + start_m + start_l = tl.load(layout_ptr).to(tl.int32) + end_l = tl.load(layout_ptr + 1).to(tl.int32) + + # Loop over k, v and update accumulator + for col_idx_idx in range(start_l, end_l): + col_idx = tl.load(layout_csr_col_indices + layout_h * layout_csr_col_stride_h + col_idx_idx).to(tl.int32) + start_n = col_idx * BLOCK_N + # -- compute qk ---- + if EVEN_N: + k = tl.load(k_ptrs + start_n * stride_kn) + else: + k = tl.load(k_ptrs + start_n * stride_kn, mask=offs_n[None, :] + start_n < total_seq_len) + qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32) + qk += tl.dot(q, k) + + if NUM_D_BLOCKS >= 2: + if EVEN_N: + k = tl.load(k_ptrs + start_n * stride_kn + BLOCK_D) + else: + k = tl.load(k_ptrs + start_n * stride_kn + BLOCK_D, mask=offs_n[None, :] + start_n < total_seq_len) + qk += tl.dot(q2, k) + + qk *= softmax_scale + + # This assumes that past sequence length is 0, otherwise need offs_m[:, None] + past_seq_len >= ... + qk += tl.where(offs_m[:, None] >= (start_n + offs_n[None, :]), 0, float("-inf")) + # -- compute m_ij, p, l_ij + m_ij = tl.max(qk, 1) + p = tl.exp(qk - m_ij[:, None]) + l_ij = tl.sum(p, 1) + # -- update m_i and l_i + m_i_new = tl.maximum(m_i, m_ij) + alpha = tl.exp(m_i - m_i_new) + beta = tl.exp(m_ij - m_i_new) + l_i_new = alpha * l_i + beta * l_ij + # -- update output accumulator -- + # scale p + p_scale = beta / l_i_new + p = p * p_scale[:, None] + # scale acc + acc_scale = l_i / l_i_new * alpha + acc = acc * acc_scale[:, None] + if NUM_D_BLOCKS >= 2: + acc2 = acc2 * acc_scale[:, None] + p = p.to(Q.dtype.element_ty) + # update acc + if EVEN_N: + v = tl.load(v_ptrs + start_n * stride_vn) + else: + v = tl.load(v_ptrs + start_n * stride_vn, mask=offs_n[:, None] + start_n < total_seq_len) + acc += tl.dot(p, v) + + if NUM_D_BLOCKS >= 2: + if EVEN_N: + v = tl.load(v_ptrs + start_n * stride_vn + BLOCK_D) + else: + v = tl.load(v_ptrs + start_n * stride_vn + BLOCK_D, mask=offs_n[:, None] + start_n < total_seq_len) + acc2 += tl.dot(p, v) + + # update m_i and l_i + l_i = l_i_new + m_i = m_i_new + + off_o = off_b * stride_ob + off_h * stride_oh + offs_m[:, None] * stride_om + offs_d[None, :] + out_ptrs = out + off_o + tl.store(out_ptrs, acc, mask=offs_m[:, None] < q_seq_len) + if NUM_D_BLOCKS >= 2: + tl.store(out_ptrs + BLOCK_D, acc2, mask=offs_m[:, None] < q_seq_len) + +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): + """ + Wrapper function to launch the Triton Block Sparse Attention kernel. + """ + q_seq_len = total_seq_len + batch_size = Q.shape[0] + + grid = (triton.cdiv(q_seq_len, BLOCK_M), batch_size * num_heads) + + out = torch.empty((batch_size, num_heads, q_seq_len, Q.shape[-1]), device=Q.device, dtype=Q.dtype) + + block_sparse_attention_kernel[grid]( + out, 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, + Q.stride(0), Q.stride(1), Q.stride(2), + K.stride(0), K.stride(1), K.stride(2), + V.stride(0), V.stride(1), V.stride(2), + out.stride(0), out.stride(1), out.stride(2), + num_heads, num_kv_heads, + total_seq_len, + BLOCK_M=BLOCK_M, EVEN_M=EVEN_M, + BLOCK_N=BLOCK_N, EVEN_N=EVEN_N, + BLOCK_D=BLOCK_D, NUM_D_BLOCKS=NUM_D_BLOCKS + ) + return out \ No newline at end of file diff --git a/benchmarks/operators/flash_attention/config.yaml b/benchmarks/operators/flash_attention/config.yaml new file mode 100644 index 00000000..38975036 --- /dev/null +++ b/benchmarks/operators/flash_attention/config.yaml @@ -0,0 +1,16 @@ +test_cases: + # Case 1: Standard LLaMA-like config (Short Context) + - batch_size: 4 + n_heads: 32 + seq_len: 1024 + head_dim: 64 + dtype: "float16" + causal: True + + # Case 2: Standard LLaMA-like config (Longer Context) + - batch_size: 1 + n_heads: 32 + seq_len: 4096 + head_dim: 64 + dtype: "float16" + causal: True diff --git a/benchmarks/operators/flash_attention/impl_cutile.py b/benchmarks/operators/flash_attention/impl_cutile.py new file mode 100644 index 00000000..d27cffb5 --- /dev/null +++ b/benchmarks/operators/flash_attention/impl_cutile.py @@ -0,0 +1,158 @@ + +import torch +import cuda.tile as ct +import math +import numpy as np +from cuda.tile import RoundingMode as RMd + +INV_LOG_2 = 1.0 / math.log(2) +ConstInt = ct.Constant[int] +ConstBool = ct.Constant[bool] + +@ct.kernel(occupancy=2) +def fmha_kernel(Q, K, V, Out, + qk_scale: float, + input_pos: int, + TILE_D: ConstInt, # TILE_D = hidden_size + H: ConstInt, + TILE_M: ConstInt, + TILE_N: ConstInt, + QUERY_GROUP_SIZE: ConstInt, + CAUSAL: ConstBool, + EVEN_K: ConstBool): + """ + cuTile kernel for Fused Multi-Head Attention (FMHA). + Computes attention output for a specific batch item and head, using tiling and online softmax. + """ + # Map block IDs to batch and head indices + bid_x = ct.bid(0) + bid_y = ct.bid(1) + batch_idx = bid_y // H + head_idx = bid_y % H + off_kv_h = head_idx // QUERY_GROUP_SIZE + + # Adjust qk_scale for exp2 + qk_scale = qk_scale * INV_LOG_2 + + # Initialize offsets for current query tile (M-dimension) + offs_m = bid_x * TILE_M + ct.arange(TILE_M, dtype=np.int32) # [TILE_M] + offs_m += input_pos + offs_m = offs_m[:, None] # [TILE_M, 1] + + # Initialize local offsets for key/value tile (N-dimension) + offs_n_tile = ct.arange(TILE_N, dtype=np.int32) # [TILE_N] + offs_n_tile = offs_n_tile[None, :] # [1, TILE_N] + + # Initialize online softmax accumulators in float32 for stability + m_i = ct.full((TILE_M, 1), -np.inf, dtype=np.float32) + l_i = ct.full((TILE_M, 1), 0.0, dtype=np.float32) + acc = ct.full((TILE_M, TILE_D), 0.0, dtype=np.float32) + + # Load query tile for this batch, head, and M-chunk + q = ct.load( + Q, index=(batch_idx, head_idx, bid_x, 0), shape=(1, 1, TILE_M, TILE_D) + ).reshape((TILE_M, TILE_D)) # [TILE_M, TILE_D] + + # loop over k, v and update accumulator + m_end = input_pos + (bid_x + 1) * TILE_M + k_seqlen = K.shape[2] + if CAUSAL: + # when kv pos could exceed q pos + mask_start = (input_pos + bid_x * TILE_M) // TILE_N + # when kv pos could exceed k_seqlen + mask_start = min(mask_start, k_seqlen // TILE_N) + Tc = ct.cdiv(min(m_end, k_seqlen), TILE_N) + else: + Tc = ct.cdiv(k_seqlen, TILE_N) + mask_start = k_seqlen // TILE_N + + # Loop over K, V blocks (N-dimension chunks) + for j in range(0, Tc): + # --- Compute QK product --- + k = ct.load( + K, index=(batch_idx, off_kv_h, 0, j), shape=(1, 1, TILE_D, TILE_N), + order=(0, 1, 3, 2), + latency=2, + ) + k = k.reshape((TILE_D, TILE_N)) # [TILE_D, TILE_N] + qk = ct.full((TILE_M, TILE_N), 0., dtype=np.float32) + qk = ct.mma(q, k, qk) # [TILE_M, TILE_N] + + # --- Apply Causal Masking --- + if (CAUSAL or not EVEN_K) and j >= mask_start: + offs_n = j * TILE_N + offs_n_tile + mask = ct.full((TILE_M, TILE_N), True, dtype=np.bool) + # out of bound mask + if not EVEN_K: + mask = mask & (offs_n < k_seqlen) + # causal mask + if CAUSAL: + mask = mask & (offs_m >= offs_n) # [TILE_M, TILE_N] + mask = ct.where(mask, 0.0, -np.inf) # [TILE_M, TILE_N] + qk += mask + + # --- Online Softmax Update --- + # Moving qk_scale multiplication after reduce_max is to improve performance. + m_ij = max(m_i, ct.max(qk, axis=-1, keepdims=True) * qk_scale) + qk = qk * qk_scale - m_ij # [TILE_M, TILE_N] + + # attention weights + p = ct.exp2(qk, flush_to_zero=True) # [TILE_M, TILE_N] + l_ij = ct.sum(p, axis=-1, keepdims=True) # [TILE_M, 1] + alpha = ct.exp2(m_i - m_ij, flush_to_zero=True) # [TILE_M, 1] + # update m_i and l_i + l_i = l_i * alpha + l_ij # [TILE_M, 1] + # scale acc + acc = acc * alpha # [TILE_M, TILE_N] + + # --- Compute PV product --- + v = ct.load( + V, index=(batch_idx, off_kv_h, j, 0), shape=(1, 1, TILE_N, TILE_D), + latency=4, + ).reshape((TILE_N, TILE_D)) # [TILE_N, TILE_D] + p = p.astype(Q.dtype) + acc = ct.mma(p, v, acc) # [TILE_M, TILE_N] + m_i = m_ij # [TILE_M, 1] + + # --- Final Normalization and Store --- + acc = ct.truediv(acc, l_i, flush_to_zero=True, rounding_mode=RMd.APPROX) + acc = acc.reshape((1, 1, TILE_M, TILE_D)).astype(Out.dtype) + ct.store(Out, index=(batch_idx, head_idx, bid_x, 0), tile=acc) + +def run(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, causal: bool = True, **kwargs): + + Batch, Heads, SeqLen_Q, D_k = q.shape + + TILE_M = 64 + TILE_N = 32 + + input_pos = 0 + + # Scale + qk_scale = 1.0 / math.sqrt(D_k) + + # EVEN_K Check + even_k = (SeqLen_Q % TILE_N) == 0 + + query_group_size = 1 + + Out = torch.empty_like(q) + + grid_x = math.ceil(SeqLen_Q / TILE_M) + grid_y = Batch * Heads + grid = (grid_x, grid_y, 1) + + ct.launch(torch.cuda.current_stream(), grid, fmha_kernel, ( + q, k, v, Out, + qk_scale, + input_pos, + D_k, + Heads, + TILE_M, + TILE_N, + query_group_size, + causal, + even_k + )) + + return Out \ No newline at end of file diff --git a/benchmarks/operators/flash_attention/impl_torch.py b/benchmarks/operators/flash_attention/impl_torch.py new file mode 100644 index 00000000..30dd4400 --- /dev/null +++ b/benchmarks/operators/flash_attention/impl_torch.py @@ -0,0 +1,5 @@ +import torch +from torch.nn.functional import scaled_dot_product_attention + +def run(q, k, v, causal=True, **kwargs): + return scaled_dot_product_attention(q, k, v, is_causal=causal) \ No newline at end of file diff --git a/benchmarks/operators/flash_attention/impl_triton.py b/benchmarks/operators/flash_attention/impl_triton.py new file mode 100644 index 00000000..6041a6d6 --- /dev/null +++ b/benchmarks/operators/flash_attention/impl_triton.py @@ -0,0 +1,122 @@ +import torch +import triton +import triton.language as tl + +@triton.jit +def _fwd_kernel( + Q, K, V, sm_scale, + L, + O, + stride_q_bs, stride_q_head, stride_q_seqlen, stride_q_dim, + stride_k_bs, stride_k_head, stride_k_seqlen, stride_k_dim, + stride_v_bs, stride_v_head, stride_v_seqlen, stride_v_dim, + stride_o_bs, stride_o_head, stride_o_seqlen, stride_o_dim, + BS, HEAD, SEQLEN, + BLOCK_M: tl.constexpr, + DIM: tl.constexpr, + BLOCK_N: tl.constexpr, + IS_CAUSAL: tl.constexpr, +): + start_m = tl.program_id(0) + off_bs_head = tl.program_id(1) + + qkv_base_offset = off_bs_head * stride_q_head + Q_block_ptr = tl.make_block_ptr( + base=Q + qkv_base_offset, + shape=(SEQLEN, DIM), + strides=(stride_q_seqlen, stride_q_dim), + offsets=(start_m * BLOCK_M, 0), + block_shape=(BLOCK_M, DIM), + order=(1, 0), + ) + K_block_ptr = tl.make_block_ptr( + base=K + qkv_base_offset, + shape=(DIM, SEQLEN), + strides=(stride_k_dim, stride_k_seqlen), + offsets=(0, 0), + block_shape=(DIM, BLOCK_N), + order=(0, 1), + ) + V_block_ptr = tl.make_block_ptr( + base=V + qkv_base_offset, + shape=(SEQLEN, DIM), + strides=(stride_k_seqlen, stride_v_dim), + offsets=(0, 0), + block_shape=(BLOCK_N, DIM), + order=(1, 0), + ) + off_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M) + off_n = tl.arange(0, BLOCK_N) + max = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf') + denom = tl.zeros([BLOCK_M], dtype=tl.float32) + out_buffer = tl.zeros([BLOCK_M, DIM], dtype=tl.float32) + qk_scale = sm_scale * 1.44269504 + q = tl.load(Q_block_ptr) + q = (q * qk_scale).to(tl.float16) + lo = 0 + hi = (start_m + 1) * BLOCK_M if IS_CAUSAL else SEQLEN + for start_n in range(lo, hi, BLOCK_N): + k = tl.load(K_block_ptr) + v = tl.load(V_block_ptr) + + qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32) + if IS_CAUSAL: + qk = tl.where(off_m[:, None] >= (start_n + off_n[None, :]), qk, float("-inf")) + qk += tl.dot(q, k) + + max_new = tl.maximum(max, tl.max(qk, 1)) + alpha = tl.math.exp2(max - max_new) + nume = tl.math.exp2(qk - max_new[:, None]) + out_scale = denom * 0 + alpha + out_buffer *= out_scale[:, None] + out_buffer += tl.dot(nume.to(tl.float16), v) + denom = denom * alpha + tl.sum(nume, 1) + max = max_new + K_block_ptr = tl.advance(K_block_ptr, (0, BLOCK_N)) + V_block_ptr = tl.advance(V_block_ptr, (BLOCK_N, 0)) + + out_buffer = out_buffer / denom[:, None] + l_ptr = L + off_bs_head * SEQLEN + off_m + tl.store(l_ptr, max + tl.math.log2(denom)) + O_block_ptr = tl.make_block_ptr( + base=O + qkv_base_offset, + shape=(SEQLEN, DIM), + strides=(stride_o_seqlen, stride_o_dim), + offsets=(start_m * BLOCK_M, 0), + block_shape=(BLOCK_M, DIM), + order=(1, 0), + ) + tl.store(O_block_ptr, out_buffer.to(tl.float16)) +def run(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, causal: bool = True, **kwargs): + + Lq, Lk, Lv = q.shape[-1], k.shape[-1], v.shape[-1] + + sm_scale = 1.0 / (Lq ** 0.5) + + o = torch.empty_like(q) + + BLOCK_M = 64 + BLOCK_N = 32 + + grid = (triton.cdiv(q.shape[2], BLOCK_M), q.shape[0] * q.shape[1], 1) + + L = torch.empty((q.shape[0] * q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32) + + num_warps = 4 if Lk <= 64 else 8 + + _fwd_kernel[grid]( + q, k, v, sm_scale, + L, + o, + q.stride(0), q.stride(1), q.stride(2), q.stride(3), + k.stride(0), k.stride(1), k.stride(2), k.stride(3), + v.stride(0), v.stride(1), v.stride(2), v.stride(3), + o.stride(0), o.stride(1), o.stride(2), o.stride(3), + q.shape[0], q.shape[1], q.shape[2], + BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, DIM=Lk, + IS_CAUSAL=causal, + num_warps=num_warps, + num_stages=4 + ) + + return o \ No newline at end of file diff --git a/benchmarks/operators/flash_decode/config.yaml b/benchmarks/operators/flash_decode/config.yaml new file mode 100644 index 00000000..9070a399 --- /dev/null +++ b/benchmarks/operators/flash_decode/config.yaml @@ -0,0 +1,24 @@ +test_cases: + - n: 1 + batch: 2 + heads: 8 + seq_len: 4096 + head_dim: 128 + block_seq: 128 + dtype: "float32" + + - n: 1 + batch: 1 + heads: 16 + seq_len: 8192 + head_dim: 128 + block_seq: 128 + dtype: "float32" + + - n: 1 + batch: 4 + heads: 32 + seq_len: 2048 + head_dim: 64 + block_seq: 64 + dtype: "float32" \ No newline at end of file diff --git a/benchmarks/operators/flash_decode/impl_cutile.py b/benchmarks/operators/flash_decode/impl_cutile.py new file mode 100644 index 00000000..f4bb97be --- /dev/null +++ b/benchmarks/operators/flash_decode/impl_cutile.py @@ -0,0 +1,85 @@ +import torch +import cuda.tile as ct +import math + +ConstInt = ct.Constant[int] + +@ct.kernel +def flash_decode_stage2_kernel( + Mid_O, # [Batch, Head, NumBlocks, HeadDim] + Mid_O_LSE, # [Batch, Head, NumBlocks] + B_Seqlen, # [Batch] + Out, # [Batch, Head, HeadDim] + HEAD_DIM: ConstInt, + BLOCK_SEQ: ConstInt, + TOTAL_BLOCKS: ConstInt +): + # 1. IDs + bid_b = ct.bid(0) + bid_h = ct.bid(1) + + + seq_len = ct.load(B_Seqlen, index=(bid_b,), shape=(1,)) + + real_num_blocks = (seq_len + BLOCK_SEQ - 1) // BLOCK_SEQ + + + acc = ct.full((1, 1, 1, HEAD_DIM), 0.0, dtype=ct.float32) + + max_logic = ct.full((1, 1, 1), -float('inf'), dtype=ct.float32) + sum_exp = ct.full((1, 1, 1), 0.0, dtype=ct.float32) + + for k in range(TOTAL_BLOCKS): + + is_valid = k < real_num_blocks + + curr_lse = ct.load(Mid_O_LSE, index=(bid_b, bid_h, k), shape=(1, 1, 1)) + + curr_o = ct.load(Mid_O, index=(bid_b, bid_h, k, 0), shape=(1, 1, 1, HEAD_DIM)) + + curr_lse = ct.where(is_valid, curr_lse, -float('inf')) + curr_o = ct.where(is_valid, curr_o, 0.0) + + new_max = ct.maximum(max_logic, curr_lse) + + scale = ct.exp(max_logic - new_max) + curr_weight = ct.exp(curr_lse - new_max) + + acc = acc * scale + + # reshape curr_weight (1,1,1) -> (1,1,1,1) + acc = acc + curr_o * curr_weight + + sum_exp = sum_exp * scale + curr_weight + max_logic = new_max + + # 5. Store + final_out = acc / sum_exp + + ct.store(Out, index=(bid_b, bid_h, 0, 0), tile=final_out) + +def run(mid_o, mid_o_lse, b_seqlen, block_seq_tensor, block_size: int = None): + + if isinstance(block_seq_tensor, torch.Tensor): + block_seq = block_seq_tensor.item() + else: + block_seq = block_seq_tensor + + batch, head_num, num_blocks, head_dim = mid_o.shape + out = torch.empty((batch, head_num, head_dim), dtype=mid_o.dtype, device=mid_o.device) + + + out_view = out.view(batch, head_num, 1, head_dim) + grid = (batch, head_num, 1) + + HEAD_DIM = head_dim + TOTAL_BLOCKS = num_blocks + + ct.launch( + torch.cuda.current_stream(), + grid, + flash_decode_stage2_kernel, + (mid_o, mid_o_lse, b_seqlen, out_view, HEAD_DIM, block_seq, TOTAL_BLOCKS) + ) + + return out \ No newline at end of file diff --git a/benchmarks/operators/flash_decode/impl_torch.py b/benchmarks/operators/flash_decode/impl_torch.py new file mode 100644 index 00000000..8af30f0a --- /dev/null +++ b/benchmarks/operators/flash_decode/impl_torch.py @@ -0,0 +1,78 @@ +import torch + +def run(mid_o, mid_o_lse, b_seqlen, block_seq, block_size=None): + """ + PyTorch reference implementation for Flash Decode Stage 2. + Performs a weighted reduction of partial attention outputs based on LogSumExp. + + Args: + mid_o: Partial outputs from Stage 1. Shape: [Batch, Heads, Num_Blocks, HeadDim] + mid_o_lse: Partial LogSumExp from Stage 1. Shape: [Batch, Heads, Num_Blocks] + b_seqlen: Actual sequence lengths. Shape: [Batch] + block_seq: Block size used in Stage 1 partitioning (scalar or 0-d Tensor). + block_size: Ignored (placeholder for benchmark framework compatibility). + + Returns: + Final output tensor of shape [Batch, Heads, HeadDim] + """ + + # Handle block_seq input (could be passed as a Tensor by the generator) + if isinstance(block_seq, torch.Tensor): + block_seq = block_seq.item() + + # Get dimensions + batch_size, num_heads, num_blocks, head_dim = mid_o.shape + + # 1. Create Validity Mask + # Calculate how many blocks are actually valid for each batch index + # Formula: ceil(seq_len / block_seq) + valid_blocks_count = (b_seqlen + block_seq - 1) // block_seq + + # Create indices for blocks [0, 1, ..., NumBlocks-1] + # Shape: [1, 1, NumBlocks] + block_indices = torch.arange(num_blocks, device=mid_o.device).view(1, 1, -1) + + # Expand valid_blocks_count to allow broadcasting + # Shape: [Batch, 1, 1] + valid_blocks_expanded = valid_blocks_count.view(-1, 1, 1) + + # Generate mask: True if the block is valid, False otherwise + # Shape: [Batch, 1, NumBlocks] -> Broadcasts to [Batch, Heads, NumBlocks] + mask = block_indices < valid_blocks_expanded + + # 2. Mask LogSumExp (LSE) + # Set LSE of invalid blocks to -inf so they don't impact the Global Max or Sum + # Clone to avoid modifying the input tensor in-place + masked_lse = mid_o_lse.clone() + masked_lse = masked_lse.masked_fill(~mask, -float('inf')) + + # 3. Compute Global Max LSE (for numerical stability) + # Find the max LSE across all blocks for each head + # Shape: [Batch, Heads, 1] + global_max_lse = torch.max(masked_lse, dim=2, keepdim=True)[0] + + # 4. Compute Weights + # Calculate exp(LSE_i - GlobalMax) + # Shape: [Batch, Heads, NumBlocks] + weights = torch.exp(masked_lse - global_max_lse) + + # Explicitly zero out weights for invalid blocks (redundant if LSE is -inf, but safe) + weights = weights.masked_fill(~mask, 0.0) + + # 5. Weighted Sum of Partial Outputs (Numerator) + # Mid_O: [Batch, Heads, NumBlocks, HeadDim] + # Weights: [Batch, Heads, NumBlocks] -> Unsqueeze to [Batch, Heads, NumBlocks, 1] + # Result: Sum across the NumBlocks dimension (dim 2) + weighted_values = mid_o * weights.unsqueeze(-1) + numerator = torch.sum(weighted_values, dim=2) + + # 6. Sum of Weights (Denominator) + # Shape: [Batch, Heads, 1] + denominator = torch.sum(weights, dim=2, keepdim=True) + + # 7. Final Normalization + # Output = Numerator / Denominator + # Add a small epsilon to denominator to prevent division by zero in empty sequences (edge case) + output = numerator / (denominator + 1e-10) + + return output \ No newline at end of file diff --git a/benchmarks/operators/flash_decode/impl_triton.py b/benchmarks/operators/flash_decode/impl_triton.py new file mode 100644 index 00000000..1656fabb --- /dev/null +++ b/benchmarks/operators/flash_decode/impl_triton.py @@ -0,0 +1,109 @@ +import torch +import triton +import triton.language as tl + +@triton.jit +def _fwd_kernel_flash_decode_stage2( + B_Seqlen, + Mid_O, # [batch, head, seq_block_num, head_dim] + Mid_O_LogExpSum, # [batch, head, seq_block_num] + Out, # [batch, head, head_dim] + stride_mid_ob, + stride_mid_oh, + stride_mid_os, + stride_mid_od, + stride_mid_o_eb, + stride_mid_o_eh, + stride_mid_o_es, + stride_obs, + stride_oh, + stride_od, + head_dim, + BLOCK_SEQ: tl.constexpr, + BLOCK_DMODEL: tl.constexpr, +): + cur_batch = tl.program_id(0) + cur_head = tl.program_id(1) + + offs_d = tl.arange(0, BLOCK_DMODEL) + cur_batch_seq_len = tl.load(B_Seqlen + cur_batch) + + block_n_size = tl.where(cur_batch_seq_len <= 0, 0, cur_batch_seq_len + BLOCK_SEQ - 1) // BLOCK_SEQ + + sum_exp = 0.0 + max_logic = -float("inf") + acc = tl.zeros([BLOCK_DMODEL], dtype=tl.float32) + + offs_v = cur_batch * stride_mid_ob + cur_head * stride_mid_oh + offs_d + offs_logic = cur_batch * stride_mid_o_eb + cur_head * stride_mid_o_eh + for block_seq_n in range(0, block_n_size, 1): + tv = tl.load(Mid_O + offs_v + block_seq_n * stride_mid_os, mask=offs_d < head_dim, other=0.0) + tlogic = tl.load(Mid_O_LogExpSum + offs_logic + block_seq_n) + new_max_logic = tl.maximum(tlogic, max_logic) + + old_scale = tl.exp(max_logic - new_max_logic) + acc *= old_scale + exp_logic = tl.exp(tlogic - new_max_logic) + acc += exp_logic * tv + sum_exp = sum_exp * old_scale + exp_logic + max_logic = new_max_logic + + tl.store(Out + cur_batch * stride_obs + cur_head * stride_oh + offs_d, acc / sum_exp, mask=offs_d < head_dim) + return + +def run(mid_o, mid_o_lse, b_seqlen, block_seq_tensor, block_size: int = None): + """ + Args: + mid_o: Partial outputs from Stage 1. Shape: [Batch, Heads, Num_Blocks, HeadDim] + mid_o_lse: Partial LogSumExp from Stage 1. Shape: [Batch, Heads, Num_Blocks] + b_seqlen: Actual sequence lengths. Shape: [Batch] + block_seq_tensor: The block size used in Stage 1 partitioning. + Can be an int or a scalar Tensor. + block_size: (Optional) Block size config from benchmark framework, usually ignored here. + """ + # Handle block_seq parameter which might be passed as a Tensor or int + if isinstance(block_seq_tensor, torch.Tensor): + block_seq = block_seq_tensor.item() + else: + block_seq = block_seq_tensor + + # Extract shapes + batch, head_num, num_blocks, head_dim = mid_o.shape + + # Allocate output tensor + output = torch.empty((batch, head_num, head_dim), dtype=mid_o.dtype, device=mid_o.device) + + # Determine Triton block size for the head dimension (must be power of 2) + BLOCK_DMODEL = triton.next_power_of_2(head_dim) + + # Grid configuration: One kernel instance per (Batch, Head) + grid = (batch, head_num) + + # Launch the kernel + _fwd_kernel_flash_decode_stage2[grid]( + B_Seqlen=b_seqlen, + Mid_O=mid_o, + Mid_O_LogExpSum=mid_o_lse, + Out=output, + # Strides for Mid_O + stride_mid_ob=mid_o.stride(0), + stride_mid_oh=mid_o.stride(1), + stride_mid_os=mid_o.stride(2), + stride_mid_od=mid_o.stride(3), + # Strides for Mid_O_LogExpSum + stride_mid_o_eb=mid_o_lse.stride(0), + stride_mid_o_eh=mid_o_lse.stride(1), + stride_mid_o_es=mid_o_lse.stride(2), + # Strides for Out + stride_obs=output.stride(0), + stride_oh=output.stride(1), + stride_od=output.stride(2), + # Constants + head_dim=head_dim, + BLOCK_SEQ=block_seq, + BLOCK_DMODEL=BLOCK_DMODEL, + num_warps=4, + num_stages=2, + ) + + return output \ No newline at end of file diff --git a/benchmarks/operators/matmul_fp16_fp8/config.yaml b/benchmarks/operators/matmul_fp16_fp8/config.yaml new file mode 100644 index 00000000..3c33e96c --- /dev/null +++ b/benchmarks/operators/matmul_fp16_fp8/config.yaml @@ -0,0 +1,9 @@ +test_cases: + - M: 1024 + K: 1024 + N: 1024 + dtype: "float16" + - M: 4096 + K: 4096 + N: 4096 + dtype: "float16" \ No newline at end of file diff --git a/benchmarks/operators/matmul_fp16_fp8/impl_cutile.py b/benchmarks/operators/matmul_fp16_fp8/impl_cutile.py new file mode 100644 index 00000000..1e39b74b --- /dev/null +++ b/benchmarks/operators/matmul_fp16_fp8/impl_cutile.py @@ -0,0 +1,77 @@ +import torch +import cuda.tile as ct +import math + +ConstInt = ct.Constant[int] + +@ct.kernel +def matmul_kernel( + A, B, C, + M: ConstInt, N: ConstInt, K: ConstInt, + TM: ConstInt, TN: ConstInt, TK: ConstInt, + GROUP_SIZE_M: ConstInt +): + """ + 1D Grid Launch with L2 Cache Swizzling. + """ + + pid = ct.bid(0) + + num_pid_m = (M + TM - 1) // TM + num_pid_n = (N + TN - 1) // TN + num_pid_in_group = GROUP_SIZE_M * num_pid_n + + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = GROUP_SIZE_M + if num_pid_m - first_pid_m < GROUP_SIZE_M: + group_size_m = num_pid_m - first_pid_m + + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + acc = ct.full((TM, TN), 0.0, ct.float32) + num_tiles_k = (K + TK - 1) // TK + for k in range(num_tiles_k): + a_tile = ct.load(A, (pid_m, k), (TM, TK), padding_mode=ct.PaddingMode.ZERO) + b_tile = ct.load(B, (k, pid_n), (TK, TN), padding_mode=ct.PaddingMode.ZERO) + acc = ct.mma(a_tile, b_tile, acc) + acc = ct.astype(acc, C.dtype) + ct.store(C, (pid_m, pid_n), acc) + + + + +def run(a: torch.Tensor, b: torch.Tensor, block_size: int = None): + """ + Wrapper function for cuTile matmul. + """ + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.dtype == b.dtype, "Incompatible dtypes" + + M, K = a.shape + _, N = b.shape + dtype = a.dtype + + c = torch.empty((M, N), device=a.device, dtype=dtype) + + if dtype == torch.float8_e4m3fn: + TM, TN, TK = 128, 256, 128 + GROUP_SIZE_M = 8 + else: # float16 + TM, TN, TK = 128, 256, 64 + GROUP_SIZE_M = 8 + + num_pid_m = math.ceil(M / TM) + num_pid_n = math.ceil(N / TN) + grid_1d = num_pid_m * num_pid_n + + grid = (grid_1d, 1, 1) + + ct.launch( + torch.cuda.current_stream(), + grid, + matmul_kernel, + (a, b, c, M, N, K, TM, TN, TK, GROUP_SIZE_M) + ) + + return c \ No newline at end of file diff --git a/benchmarks/operators/matmul_fp16_fp8/impl_torch.py b/benchmarks/operators/matmul_fp16_fp8/impl_torch.py new file mode 100644 index 00000000..56403c90 --- /dev/null +++ b/benchmarks/operators/matmul_fp16_fp8/impl_torch.py @@ -0,0 +1,19 @@ +import torch + +def run(a, b): + """ + Reference implementation for Matrix Multiplication. + Handles standard types and provides a safe fallback for FP8. + """ + orig_dtype = a.dtype + + # PyTorch's native `torch.matmul` may require explicit scaling functions + # (e.g., `_scaled_mm`) for FP8 tensors depending on the version. + # To match the Triton kernel's behavior (which accumulates in FP32 and casts back), + # we compute the reference in FP32 and cast the result back. + if orig_dtype in (torch.float8_e4m3fn, torch.float8_e5m2): + out = torch.matmul(a.to(torch.float32), b.to(torch.float32)) + return out.to(orig_dtype) + + # For float16, bfloat16, float32, etc. + return torch.matmul(a, b) \ No newline at end of file diff --git a/benchmarks/operators/matmul_fp16_fp8/impl_triton.py b/benchmarks/operators/matmul_fp16_fp8/impl_triton.py new file mode 100644 index 00000000..dda3de62 --- /dev/null +++ b/benchmarks/operators/matmul_fp16_fp8/impl_triton.py @@ -0,0 +1,128 @@ +import torch +import triton +import triton.language as tl + +def _matmul_launch_metadata(grid, kernel, args): + ret = {} + M, N, K = args["M"], args["N"], args["K"] + ret["name"] = f"{kernel.name} [M={M}, N={N}, K={K}]" + if "c_ptr" in args: + bytes_per_elem = args["c_ptr"].element_size() + else: + bytes_per_elem = 1 if args["FP8_OUTPUT"] else 2 + ret[f"flops{bytes_per_elem * 8}"] = 2. * M * N * K + ret["bytes"] = bytes_per_elem * (M * K + N * K + M * N) + return ret + + +@triton.jit(launch_metadata=_matmul_launch_metadata) +def matmul_kernel(a_ptr, b_ptr, c_ptr, # + M, N, K, # + stride_am, stride_ak, # + stride_bk, stride_bn, # + stride_cm, stride_cn, # + BLOCK_SIZE_M: tl.constexpr, # + BLOCK_SIZE_N: tl.constexpr, # + BLOCK_SIZE_K: tl.constexpr, # + GROUP_SIZE_M: tl.constexpr, # + ): + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + start_m = pid_m * BLOCK_SIZE_M + start_n = pid_n * BLOCK_SIZE_N + + offs_am = start_m + tl.arange(0, BLOCK_SIZE_M) + offs_bn = start_n + tl.arange(0, BLOCK_SIZE_N) + offs_am = tl.where(offs_am < M, offs_am, 0) + offs_bn = tl.where(offs_bn < N, offs_bn, 0) + + offs_am = tl.max_contiguous(tl.multiple_of(offs_am, BLOCK_SIZE_M), BLOCK_SIZE_M) + offs_bn = tl.max_contiguous(tl.multiple_of(offs_bn, BLOCK_SIZE_N), BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + + for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) + accumulator = tl.dot(a, b, accumulator) + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + + if (c_ptr.dtype.element_ty == tl.float8e4nv): + c = accumulator.to(tl.float8e4nv) + else: + c = accumulator.to(tl.float16) + + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, c, mask=c_mask) + +# ============================================================================== +# END OF TRITON SOURCE KERNEL +# ============================================================================== + +def run(a: torch.Tensor, b: torch.Tensor, block_size: int = None): + """ + Wrapper function to execute the Triton matmul kernel. + The benchmark framework will call this function. + """ + # Validation + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + assert a.dtype == b.dtype, "Incompatible dtypes" + + M, K = a.shape + _, N = b.shape + dtype = a.dtype + + # Configuration map matching the original source code + configs = { + torch.float8_e4m3fn: { + "BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 128, "GROUP_SIZE_M": 8, "num_stages": 4, + "num_warps": 8 + }, + torch.float16: { + "BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 64, "GROUP_SIZE_M": 8, "num_stages": 3, + "num_warps": 8 + } + } + + if dtype not in configs: + raise ValueError(f"Dtype {dtype} is not configured for this Triton kernel.") + + config = configs[dtype] + + # Allocate output tensor + c = torch.empty((M, N), device=a.device, dtype=dtype) + + # Grid function + grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]), ) + + # Launch kernel + matmul_kernel[grid]( + a, b, c, + M, N, K, + a.stride(0), a.stride(1), + b.stride(0), b.stride(1), + c.stride(0), c.stride(1), + BLOCK_SIZE_M=config["BLOCK_SIZE_M"], + BLOCK_SIZE_N=config["BLOCK_SIZE_N"], + BLOCK_SIZE_K=config["BLOCK_SIZE_K"], + GROUP_SIZE_M=config["GROUP_SIZE_M"], + num_stages=config["num_stages"], + num_warps=config["num_warps"], + ) + + return c \ No newline at end of file diff --git a/benchmarks/operators/matmul_int8/config.yaml b/benchmarks/operators/matmul_int8/config.yaml new file mode 100644 index 00000000..15b41653 --- /dev/null +++ b/benchmarks/operators/matmul_int8/config.yaml @@ -0,0 +1,20 @@ +test_cases: + - M: 1024 + K_b: 256 + N: 1024 + block_size: 128 + + - M: 4096 + K_b: 1024 + N: 4096 + block_size: 128 + + - M: 8192 + K_b: 2048 + N: 8192 + block_size: 128 + + - M: 2048 + K_b: 1024 + N: 8192 + block_size: 128 \ No newline at end of file diff --git a/benchmarks/operators/matmul_int8/impl_cutile.py b/benchmarks/operators/matmul_int8/impl_cutile.py new file mode 100644 index 00000000..a479f642 --- /dev/null +++ b/benchmarks/operators/matmul_int8/impl_cutile.py @@ -0,0 +1,75 @@ +import torch +import cuda.tile as ct +import math + +ConstInt = ct.Constant[int] + +@ct.kernel +def matmul_int8_kernel( + A, B, C, + TM: ConstInt, TN: ConstInt, TK: ConstInt, + GROUP_SIZE_M: ConstInt +): + + M = A.shape[0] + K = A.shape[1] + N = B.shape[1] + K_b = B.shape[0] + + pid = ct.bid(0) + num_pid_m = ct.cdiv(M, TM) + num_pid_n = ct.cdiv(N, TN) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = ct.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M) + + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + acc = ct.full((TM, TN), 0, ct.int32) + + num_tiles_kb = ct.cdiv(K_b, TK) + for i in range(4): + for j in range(num_tiles_kb): + k_a = i * num_tiles_kb + j + A_tile = ct.astype(ct.load(A, (pid_m, k_a), (TM, TK), padding_mode=ct.PaddingMode.ZERO), ct.int8) + B_tile = ct.astype(ct.load(B, (j, pid_n), (TK, TN), padding_mode=ct.PaddingMode.ZERO), ct.int8) + mask = 3 << (2 * i) + B_tile = ct.astype(B_tile, ct.int32) + b_unpacked = ct.astype(((B_tile & mask) >> (2 * i)), ct.int8) + b_unpacked = b_unpacked - 1 + acc = ct.mma(A_tile, b_unpacked, acc) + + ct.store(C, index=(pid_m, pid_n), tile=acc) + + +def run(a: torch.Tensor, b: torch.Tensor, block_size: int = None): + """ + Wrapper function for cuTile matmul_int8. + """ + assert a.shape[1] == b.shape[0] * 4, "Incompatible dimensions" + + M, K = a.shape + K_b, N = b.shape + + c = torch.empty((M, N), device=a.device, dtype=torch.int32) + + TM, TN, TK = 128, 128, 32 + GROUP_SIZE_M = 8 + + num_pid_m = math.ceil(M / TM) + num_pid_n = math.ceil(N / TN) + grid_1d = num_pid_m * num_pid_n + + grid = (grid_1d, 1, 1) + + ct.launch( + torch.cuda.current_stream(), + grid, + matmul_int8_kernel, + (a, b, c, TM, TN, TK, GROUP_SIZE_M) + ) + + return c \ No newline at end of file diff --git a/benchmarks/operators/matmul_int8/impl_torch.py b/benchmarks/operators/matmul_int8/impl_torch.py new file mode 100644 index 00000000..81d4844f --- /dev/null +++ b/benchmarks/operators/matmul_int8/impl_torch.py @@ -0,0 +1,21 @@ +import torch + +def run(a, b): + a_int8 = a.to(torch.int8) + + M, K = a.shape + K_b, N = b.shape + + b_unpacked = torch.empty((K, N), dtype=torch.int8, device=a.device) + for i in range(4): + mask = 3 << (2 * i) + + b_val = ((b.to(torch.int32) & mask) >> (2 * i)).to(torch.int8) + + b_val = b_val - 1 + + b_unpacked[i * K_b : (i + 1) * K_b, :] = b_val + + out = torch.matmul(a_int8.to(torch.float32), b_unpacked.to(torch.float32)) + + return out.to(torch.int32) \ No newline at end of file diff --git a/benchmarks/operators/matmul_int8/impl_triton.py b/benchmarks/operators/matmul_int8/impl_triton.py new file mode 100644 index 00000000..ea3d15f9 --- /dev/null +++ b/benchmarks/operators/matmul_int8/impl_triton.py @@ -0,0 +1,105 @@ +import torch +import triton +import triton.language as tl + +@triton.jit +def matmul_kernel( + a_ptr, + b_ptr, + c_ptr, + M, + N, + K: tl.constexpr, + stride_am, + stride_ak, + stride_bk, + stride_bn, + stride_cm, + stride_cn, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + GROUP_SIZE_M: tl.constexpr, +): + tl.static_assert( + K % (4 * BLOCK_SIZE_K) == 0, + "K / 4 must be divisible by BLOCK_SIZE_K => K divisible by 4*BLOCK_SIZE_K", + ) + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.int32) + for i in range(4): + b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + for j in range(0, tl.cdiv(K // 4, BLOCK_SIZE_K)): + k = i * tl.cdiv(K // 4, BLOCK_SIZE_K) + j + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0).to(tl.int8) + b_uint8 = tl.load(b_ptrs, mask=offs_k[:, None] < K, other=0) + mask = 3 << (2 * i) + b = ((b_uint8 & mask) >> (2 * i)).to(tl.int8) + tensor_full = tl.full((1,), 1, dtype=tl.int8) + accumulator += tl.dot(a, (b - tensor_full), out_dtype=tl.int32) + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + c = accumulator + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, c, mask=c_mask) + +def run(a: torch.Tensor, b: torch.Tensor, block_size: int = None): + + assert a.shape[1] == b.shape[0] * 4, "Incompatible dimensions, the weight matrix need to be packed" + assert a.is_contiguous(), "Matrix A must be contiguous" + + M, K = a.shape + _, N = b.shape + + + c = torch.empty((M, N), device=a.device, dtype=torch.int32) + + BLOCK_M = 128 + BLOCK_N = 128 + BLOCK_K = 64 + GROUP_M = 8 + num_stages = 4 + num_warps = 4 + + grid = lambda META: ( + triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) + + matmul_kernel[grid]( + a, + b, + c, + M, + N, + K, + a.stride(0), + a.stride(1), + b.stride(0), + b.stride(1), + c.stride(0), + c.stride(1), + BLOCK_SIZE_M=BLOCK_M, + BLOCK_SIZE_N=BLOCK_N, + BLOCK_SIZE_K=BLOCK_K, + GROUP_SIZE_M=GROUP_M, + num_stages=num_stages, + num_warps=num_warps, + ) + + return c \ No newline at end of file diff --git a/benchmarks/operators/rope/config.yaml b/benchmarks/operators/rope/config.yaml new file mode 100644 index 00000000..d3814e41 --- /dev/null +++ b/benchmarks/operators/rope/config.yaml @@ -0,0 +1,37 @@ +test_cases: + # Total: 1 * 4096 * 32 * 128 = 16MB elements + - batch_size: 1 + seq_len: 4096 + n_heads: 32 + head_dim: 128 + dtype: "float32" + + # Case 2: High Throughput / Server Scenario (Large Batch) + # Total: 16 * 512 * 32 * 128 = 32MB elements + - batch_size: 16 + seq_len: 512 + n_heads: 32 + head_dim: 128 + dtype: "float32" + + # Case 3: Smaller Head Dim (e.g., older models or specific archs) + # Total: 4 * 2048 * 16 * 64 = 8MB elements + - batch_size: 4 + seq_len: 2048 + n_heads: 16 + head_dim: 64 + dtype: "float32" + + # Case 4: Heavy Load (Large Batch + Mid Seq) + # Total: 8 * 1024 * 32 * 128 = 32MB elements + - batch_size: 8 + seq_len: 1024 + n_heads: 32 + head_dim: 128 + dtype: "float32" + + - batch_size: 8 + seq_len: 1024 + n_heads: 32 + head_dim: 128 + dtype: "float16" \ No newline at end of file diff --git a/benchmarks/operators/rope/impl_cutile.py b/benchmarks/operators/rope/impl_cutile.py new file mode 100644 index 00000000..74f34a40 --- /dev/null +++ b/benchmarks/operators/rope/impl_cutile.py @@ -0,0 +1,56 @@ +import torch +import cuda.tile as ct +import math + + +ConstInt = ct.Constant[int] + +@ct.kernel +def rope_kernel( + Q, # Rank 4: [TotalTokens, Heads, 2, HalfDim] + Cos, # Rank 2: [SeqLen, HalfDim] + Sin, # Rank 2: [SeqLen, HalfDim] + SeqLen: ConstInt, + TILE_DIM: ConstInt # HalfDim +): + + row_id = ct.bid(0) # Batch*Seq + head_id = ct.bid(1) # Heads + + seq_idx = row_id % SeqLen + + cos_tile = ct.load(Cos, index=(seq_idx, 0), shape=(1, TILE_DIM)) + sin_tile = ct.load(Sin, index=(seq_idx, 0), shape=(1, TILE_DIM)) + + q1 = ct.load(Q, index=(row_id, head_id, 0, 0), shape=(1, 1, 1, TILE_DIM)) + + q2 = ct.load(Q, index=(row_id, head_id, 1, 0), shape=(1, 1, 1, TILE_DIM)) + + out1 = q1 * cos_tile - q2 * sin_tile + out2 = q2 * cos_tile + q1 * sin_tile + + ct.store(Q, index=(row_id, head_id, 0, 0), tile=out1) + ct.store(Q, index=(row_id, head_id, 1, 0), tile=out2) + + +def run(q: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, block_size: int = None): + output = q.clone().contiguous() + + batch, seq_len, n_heads, head_dim = output.shape + half_dim = head_dim // 2 + + output_view = output.view(batch * seq_len, n_heads, 2, half_dim) + + cos_view = cos.view(seq_len, half_dim) + sin_view = sin.view(seq_len, half_dim) + + grid = (batch * seq_len, n_heads, 1) + + ct.launch( + torch.cuda.current_stream(), + grid, + rope_kernel, + (output_view, cos_view, sin_view, seq_len, half_dim) + ) + + return output \ No newline at end of file diff --git a/benchmarks/operators/rope/impl_torch.py b/benchmarks/operators/rope/impl_torch.py new file mode 100644 index 00000000..bb381597 --- /dev/null +++ b/benchmarks/operators/rope/impl_torch.py @@ -0,0 +1,23 @@ +import torch + +def rotate_half(x): + """Rotates half the hidden dims of the input.""" + x1 = x[..., : x.shape[-1] // 2] + x2 = x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + +def run(q, cos, sin): + # q: [batch, seq_len, n_heads, head_dim] + # cos, sin: [seq_len, head_dim // 2] + + cos = cos.unsqueeze(0).unsqueeze(2) # [1, seq_len, 1, head_dim/2] + sin = sin.unsqueeze(0).unsqueeze(2) + + head_dim = q.shape[-1] + q1 = q[..., :head_dim//2] + q2 = q[..., head_dim//2:] + + q1_out = (q1 * cos) - (q2 * sin) + q2_out = (q2 * cos) + (q1 * sin) + + return torch.cat((q1_out, q2_out), dim=-1) \ No newline at end of file diff --git a/benchmarks/operators/rope/impl_triton.py b/benchmarks/operators/rope/impl_triton.py new file mode 100644 index 00000000..8845205c --- /dev/null +++ b/benchmarks/operators/rope/impl_triton.py @@ -0,0 +1,94 @@ +import torch +import triton +import triton.language as tl + +MAX_FUSED_SIZE = 65536 +next_power_of_2 = triton.next_power_of_2 + +def calculate_settings(n): + BLOCK_SIZE = next_power_of_2(n) + if BLOCK_SIZE > MAX_FUSED_SIZE: + raise RuntimeError(f"Cannot launch Triton kernel since n = {n} exceeds "\ + f"the maximum CUDA blocksize = {MAX_FUSED_SIZE}.") + num_warps = 4 + if BLOCK_SIZE >= 32768: num_warps = 32 + elif BLOCK_SIZE >= 8192: num_warps = 16 + elif BLOCK_SIZE >= 2048: num_warps = 8 + return BLOCK_SIZE, num_warps + +ROPE_GROUP_SIZE = 4 + +@triton.heuristics({"BACKWARD_PASS": lambda args: args["BACKWARD_PASS"],}) +@triton.jit +def _rope_embedding( + Q, Q_row_stride, + cos, cos_row_stride, + sin, sin_row_stride, + seqlen, + head_dim : tl.constexpr, + n_heads : tl.constexpr, + BACKWARD_PASS : tl.constexpr, + BLOCK_SIZE : tl.constexpr, + ROPE_GROUP_SIZE : tl.constexpr = 4, +): + """ + Calculates the RoPE Embedding quickly + RoPE is Q * cos + rotate_half(Q) * sin + See our blog post for more info + """ + row_position = tl.program_id(0) + group_head_position = tl.program_id(1) + col_offsets = tl.arange(0, BLOCK_SIZE) + half_head_dim = head_dim // 2 + mask = col_offsets < half_head_dim + + sin1 = tl.load(sin + (row_position % seqlen)*sin_row_stride + \ + half_head_dim*0 + col_offsets, mask = mask, other = 0) + cos1 = tl.load(cos + (row_position % seqlen)*cos_row_stride + \ + half_head_dim*0 + col_offsets, mask = mask, other = 0) + + if BACKWARD_PASS: + # See our blog post for more info. + sin1 = -sin1 + + # [TODO] Autotune ROPE_GROUP_SIZE to be 1, 2, 4, 8 + head_start = group_head_position * ROPE_GROUP_SIZE + head_end = min((head_start + ROPE_GROUP_SIZE), n_heads) + + # 10% Faster kernel from [HuyNguyen-hust](https://github.com/unslothai/unsloth/pull/238) + for k in range(head_start, head_end): + offs_q1 = row_position * Q_row_stride + k * head_dim + col_offsets + offs_q2 = row_position * Q_row_stride + k * head_dim + col_offsets + half_head_dim + + # For Gemma - sometimes RoPE must be done in float32 and not bfloat16 + Q1 = tl.load(Q + offs_q1, mask = mask, other = 0).to(sin1.dtype) + Q2 = tl.load(Q + offs_q2, mask = mask, other = 0).to(sin1.dtype) + + tl.store(Q + offs_q1, Q1*cos1 - Q2*sin1, mask = mask) + tl.store(Q + offs_q2, Q2*cos1 + Q1*sin1, mask = mask) + +def run(q: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, block_size: int = None): + + output = q.clone().contiguous() + + batch, seq_len, n_heads, head_dim = output.shape + + BLOCK_SIZE, num_warps = calculate_settings(head_dim // 2) + + n_rows = batch * seq_len + div, mod = divmod(n_heads, ROPE_GROUP_SIZE) + n_groups = div + (mod != 0) + + _rope_embedding[(n_rows, n_groups, )]( + output, output.stride(1), + cos, cos.stride(0), + sin, sin.stride(0), + seq_len, + head_dim, + n_heads, + BACKWARD_PASS = False, + BLOCK_SIZE = BLOCK_SIZE, + num_warps = num_warps, + ) + + return output \ No newline at end of file diff --git a/benchmarks/operators/softmax/config.yaml b/benchmarks/operators/softmax/config.yaml new file mode 100644 index 00000000..342a6e15 --- /dev/null +++ b/benchmarks/operators/softmax/config.yaml @@ -0,0 +1,20 @@ +test_cases: + - n: 1048576 # 1024 * 1024 + shape: [1024, 1024] + block_size: 1024 + dtype: "float32" + + - n: 1024000 + shape: [1024, 1000] + block_size: 1024 + dtype: "float32" + + - n: 528384 # 4096 * 129 + shape: [4096, 129] + block_size: 256 + dtype: "float32" + + - n: 8388608 # 2048 * 4096 + shape: [2048, 4096] + block_size: 4096 + dtype: "float32" \ No newline at end of file diff --git a/benchmarks/operators/softmax/impl_cutile.py b/benchmarks/operators/softmax/impl_cutile.py new file mode 100644 index 00000000..166b1bbb --- /dev/null +++ b/benchmarks/operators/softmax/impl_cutile.py @@ -0,0 +1,53 @@ +import torch +import cuda.tile as ct +import math + + +ConstInt = ct.Constant[int] + +@ct.kernel +def softmax_kernel( + input_tensor, + output_tensor, + N_COLS: ConstInt, + TILE_SIZE: ConstInt +): + """ + input_tensor: (Rows, Cols) + output_tensor: (Rows, Cols) + """ + row_idx = ct.bid(0) + + tile = ct.load(input_tensor, index=(row_idx, 0), shape=(1, TILE_SIZE), padding_mode=ct.PaddingMode.NEG_INF) + + max_val = ct.max(tile) + + exp_tile = ct.exp(tile - max_val) + + sum_val = ct.sum(exp_tile) + + output_tile = exp_tile / sum_val + ct.store(output_tensor, index=(row_idx, 0), tile=output_tile) + + +def run(x: torch.Tensor, block_size: int): + n_rows, n_cols = x.shape + + + if block_size < n_cols: + raise RuntimeError(f"Block size ({block_size}) must be >= n_cols ({n_cols}) for cuTile Softmax.") + + output = torch.empty_like(x) + + + grid = (n_rows, 1, 1) + + + ct.launch( + torch.cuda.current_stream(), + grid, + softmax_kernel, + (x, output, n_cols, block_size) + ) + + return output \ No newline at end of file diff --git a/benchmarks/operators/softmax/impl_torch.py b/benchmarks/operators/softmax/impl_torch.py new file mode 100644 index 00000000..da0b2f0d --- /dev/null +++ b/benchmarks/operators/softmax/impl_torch.py @@ -0,0 +1,4 @@ +import torch + +def run(x): + return torch.softmax(x, dim=-1) \ No newline at end of file diff --git a/benchmarks/operators/softmax/impl_triton.py b/benchmarks/operators/softmax/impl_triton.py new file mode 100644 index 00000000..8c2c63a3 --- /dev/null +++ b/benchmarks/operators/softmax/impl_triton.py @@ -0,0 +1,53 @@ +import torch +import triton +import triton.language as tl +@triton.jit +def softmax_kernel( + output_ptr, input_ptr, input_row_stride, output_row_stride, n_cols, + BLOCK_SIZE: tl.constexpr, +): + # Get the program ID + row_idx = tl.program_id(axis=0) + + # Compute the memory offsets for this row + row_start_ptr = input_ptr + row_idx * input_row_stride + out_row_start_ptr = output_ptr + row_idx * output_row_stride + + # Load the row into SRAM + row = tl.load(row_start_ptr + tl.arange(0, BLOCK_SIZE), mask=tl.arange(0, BLOCK_SIZE) < n_cols, other=-float('inf')) + + # Compute max for numerical stability + row_max = tl.max(row, axis=0) + + # Subtract max from row and exponentiate + numerator = tl.exp(row - row_max) + + # Compute sum for normalization + denominator = tl.sum(numerator, axis=0) + + # Normalize + softmax_output = numerator / denominator + + # Store the output + tl.store(out_row_start_ptr + tl.arange(0, BLOCK_SIZE), softmax_output, mask=tl.arange(0, BLOCK_SIZE) < n_cols) + +def run(x: torch.Tensor, block_size: int): + n_rows, n_cols = x.shape + output = torch.empty_like(x) + + if block_size < n_cols: + raise RuntimeError(f"Block size ({block_size}) must be >= n_cols ({n_cols}) for this Softmax kernel.") + + if (block_size & (block_size - 1)) != 0: + raise RuntimeError(f"Block size ({block_size}) must be a power of 2.") + + grid = (n_rows,) + + softmax_kernel[grid]( + output, x, + x.stride(0), output.stride(0), + n_cols, + BLOCK_SIZE=block_size + ) + + return output \ No newline at end of file diff --git a/benchmarks/operators/streamk_matmul/config.yaml b/benchmarks/operators/streamk_matmul/config.yaml new file mode 100644 index 00000000..133f9965 --- /dev/null +++ b/benchmarks/operators/streamk_matmul/config.yaml @@ -0,0 +1,29 @@ +test_cases: + + - M: 2048 + K: 2048 + N: 2048 + dtype: "float32" + block_size: 128 + grid_programs: 108 + + - M: 2176 + K: 2176 + N: 2176 + dtype: "float32" + block_size: 128 + grid_programs: 108 + + - M: 256 + K: 8192 + N: 4096 + dtype: "float32" + block_size: 128 + grid_programs: 108 + + - M: 4096 + K: 8192 + N: 5120 + dtype: "float32" + block_size: 128 + grid_programs: 108 \ No newline at end of file diff --git a/benchmarks/operators/streamk_matmul/impl_cutile.py b/benchmarks/operators/streamk_matmul/impl_cutile.py new file mode 100644 index 00000000..910da1b4 --- /dev/null +++ b/benchmarks/operators/streamk_matmul/impl_cutile.py @@ -0,0 +1,217 @@ +import torch +import cuda.tile as ct +import math + +ConstInt = ct.Constant[int] + +# ============================================================================== +# Helper Functions (Inlined by cuTile trace) +# ============================================================================== +def get_swizzle_tile_coords(tile_id, M, N, TM, TN, GROUP_M): + """ + Translates a linear tile_id into a 2D (pid_m, pid_n) coordinate + using L2 cache swizzling. + """ + grid_m = (M + TM - 1) // TM + grid_n = (N + TN - 1) // TN + width = GROUP_M * grid_n + + group_id = tile_id // width + group_size = ct.minimum(grid_m - group_id * GROUP_M, GROUP_M) + + pid_m = group_id * GROUP_M + (tile_id % group_size) + pid_n = (tile_id % width) // group_size + return pid_m, pid_n + +# ============================================================================== +# Kernels +# ============================================================================== + +@ct.kernel +def first_wave_kernel( + A, B, C, Locks, + TM: ConstInt, TN: ConstInt, TK: ConstInt, + total_full_tiles_streamk: ConstInt, + total_partial_tiles_streamk: ConstInt, + iters_per_tile: ConstInt, + GROUP_M: ConstInt +): + """ + The Stream-K worker kernel. It computes an assigned range of K-iterations. + """ + # Fetch dynamic shapes to prevent loop unrolling + M = A.shape[0] + N = B.shape[1] + K = A.shape[1] + + pid = ct.bid(0) + + # Calculate the global iteration range assigned to this specific block (SM) + start_iter = pid * total_full_tiles_streamk + ct.minimum(pid, total_partial_tiles_streamk) + last_iter = (pid + 1) * total_full_tiles_streamk + ct.minimum(pid + 1, total_partial_tiles_streamk) + + # Loop over the assigned iteration chunks + while start_iter < last_iter: + # Determine where this tile's iterations end + rem = iters_per_tile - (start_iter % iters_per_tile) + end_iter = ct.minimum(start_iter + rem, last_iter) + + # Calculate Tile coordinates + tile_id = start_iter // iters_per_tile + pid_m, pid_n = get_swizzle_tile_coords(tile_id, M, N, TM, TN, GROUP_M) + + # Initialize FP32 accumulator for MMA + acc = ct.full((TM, TN), 0.0, dtype=ct.float32) + + # MMA Loop (Using while to avoid compile-time loop unrolling) + current_iter = start_iter + while current_iter < end_iter: + k_chunk = current_iter % iters_per_tile + + # Load tiles safely with padding + a_tile = ct.load(A, index=(pid_m, k_chunk), shape=(TM, TK), padding_mode=ct.PaddingMode.ZERO) + b_tile = ct.load(B, index=(k_chunk, pid_n), shape=(TK, TN), padding_mode=ct.PaddingMode.ZERO) + + # Tensor Core MMA + acc = ct.mma(a_tile, b_tile, acc) + current_iter = current_iter + 1 + + # Downcast accumulator to target dtype + acc_casted = ct.astype(acc, C.dtype) + + # Write-back and Synchronization Logic + if end_iter % iters_per_tile == 0: + # 1. This block computed the END of the tile. + # We can safely store to C. + ct.store(C, index=(pid_m, pid_n), tile=acc_casted) + + # If we didn't start at the beginning, we must unlock the spin-lock + # for the block that computed the partial start. + if start_iter % iters_per_tile != 0: + ct.atomic_xchg(Locks, (tile_id,), 1) + else: + # 2. This block computed the START of a tile but didn't finish it. + # Must wait for the finishing block to write C first. + while ct.atomic_cas(Locks, (tile_id,), 1, 1) != 1: + pass + + # Construct broadcastable indices for Bulk Atomic Add + # row_indices: (TM, 1) + row_indices = ct.expand_dims(pid_m * TM + ct.arange(TM, dtype=ct.int32), 1) + # col_indices: (1, TN) + col_indices = ct.expand_dims(pid_n * TN + ct.arange(TN, dtype=ct.int32), 0) + + # Bulk atomic add (cuTile handles OOB indices automatically) + ct.atomic_add(C, (row_indices, col_indices), acc_casted) + + # Move to the next chunk + start_iter = end_iter + + +@ct.kernel +def full_tiles_kernel( + A, B, C, + TM: ConstInt, TN: ConstInt, TK: ConstInt, + total_tiles_streamk: ConstInt, + GROUP_M: ConstInt +): + """ + Standard Data-Parallel GEMM for leftover full tiles. + """ + M = A.shape[0] + N = B.shape[1] + K = A.shape[1] + + pid_offset = ct.bid(0) + tile_id = pid_offset + total_tiles_streamk + + # 1. Map tile_id to (pid_m, pid_n) + pid_m, pid_n = get_swizzle_tile_coords(tile_id, M, N, TM, TN, GROUP_M) + + # 2. Accumulator + acc = ct.full((TM, TN), 0.0, dtype=ct.float32) + + # 3. K-dimension Loop (Using while to avoid unrolling) + num_tiles_k = (K + TK - 1) // TK + k = 0 + while k < num_tiles_k: + a_tile = ct.load(A, index=(pid_m, k), shape=(TM, TK), padding_mode=ct.PaddingMode.ZERO) + b_tile = ct.load(B, index=(k, pid_n), shape=(TK, TN), padding_mode=ct.PaddingMode.ZERO) + acc = ct.mma(a_tile, b_tile, acc) + k = k + 1 + + # 4. Write back + acc_casted = ct.astype(acc, C.dtype) + ct.store(C, index=(pid_m, pid_n), tile=acc_casted) + +# ============================================================================== +# Runner Wrapper +# ============================================================================== +def run(a: torch.Tensor, b: torch.Tensor, block_size: int = None, + grid_programs: int = 16, # Adjust based on target GPU SM count (e.g. 108 for A100) + BLK_M: int = 128, BLK_N: int = 128, BLK_K: int = 32, + two_tiles: bool = True): + + assert a.shape[1] == b.shape[0], "Incompatible dimensions" + + M, K = a.shape + _, N = b.shape + dtype = a.dtype + + # Configure shapes based on dtype + if dtype == torch.float8_e4m3fn: + TM, TN, TK = 128, 256, 128 + else: + TM, TN, TK = BLK_M, BLK_N, BLK_K + + GROUP_M = 8 + + # 1. CPU-side Scheduler Math + total_blocks_M = math.ceil(M / TM) + total_blocks_N = math.ceil(N / TN) + iters_per_tile = math.ceil(K / TK) + + total_tiles = total_blocks_M * total_blocks_N + total_programs_streamk = grid_programs + + if total_programs_streamk > 0: + total_tiles_streamk = total_tiles % total_programs_streamk + if two_tiles and total_tiles - total_tiles_streamk > total_programs_streamk: + total_tiles_streamk += total_programs_streamk + + total_blocking_tiles = total_tiles - total_tiles_streamk + total_iters_streamk = total_tiles_streamk * iters_per_tile + total_full_tiles_streamk = total_iters_streamk // total_programs_streamk + total_partial_tiles_streamk = total_iters_streamk % total_programs_streamk + else: + total_blocking_tiles = total_tiles + total_tiles_streamk = 0 + total_full_tiles_streamk = 0 + total_partial_tiles_streamk = 0 + + # 2. Allocate outputs and locks + c = torch.empty((M, N), device=a.device, dtype=dtype) + # Ensure locks are initialized to 0 + locks = torch.zeros((max(1, total_tiles_streamk),), device=a.device, dtype=torch.int32) + + stream = torch.cuda.current_stream() + + # 3. Launch first_wave_kernel (Stream-K) + if total_programs_streamk > 0: + grid_1 = (total_programs_streamk, 1, 1) + ct.launch( + stream, grid_1, first_wave_kernel, + (a, b, c, locks, TM, TN, TK, + total_full_tiles_streamk, total_partial_tiles_streamk, + iters_per_tile, GROUP_M) + ) + + # 4. Launch full_tiles_kernel (Data Parallel) + if total_blocking_tiles > 0: + grid_2 = (total_blocking_tiles, 1, 1) + ct.launch( + stream, grid_2, full_tiles_kernel, + (a, b, c, TM, TN, TK, total_tiles_streamk, GROUP_M) + ) + + return c \ No newline at end of file diff --git a/benchmarks/operators/streamk_matmul/impl_torch.py b/benchmarks/operators/streamk_matmul/impl_torch.py new file mode 100644 index 00000000..1ab8fae7 --- /dev/null +++ b/benchmarks/operators/streamk_matmul/impl_torch.py @@ -0,0 +1,14 @@ +import torch + +def run(a, b, **kwargs): + """ + Reference implementation for Matrix Multiplication (Stream-K target). + """ + orig_dtype = a.dtype + + # Safe handling for FP8 if necessary + if orig_dtype in (torch.float8_e4m3fn, torch.float8_e5m2): + out = torch.matmul(a.to(torch.float32), b.to(torch.float32)) + return out.to(orig_dtype) + + return torch.matmul(a, b) \ No newline at end of file diff --git a/benchmarks/operators/streamk_matmul/impl_triton.py b/benchmarks/operators/streamk_matmul/impl_triton.py new file mode 100644 index 00000000..30ca7786 --- /dev/null +++ b/benchmarks/operators/streamk_matmul/impl_triton.py @@ -0,0 +1,206 @@ +import torch +import triton +from triton import language as tl + +@triton.jit() +def swizzle_tile(tile_id, + M, N, K, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr + ): + grid_m = tl.cdiv(M, BLOCK_M) + grid_n = tl.cdiv(N, BLOCK_N) + # re-order program ID for better L2 performance + width = GROUP_M * grid_n + group_id = tile_id // width + group_size = tl.minimum(grid_m - group_id * GROUP_M, GROUP_M) + pid_m = group_id * GROUP_M + (tile_id % group_size) + pid_n = (tile_id % width) // group_size + return pid_m, pid_n + +@triton.jit() +def linear_tile(tile_id, + M, N, K, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr + ): + pid_m = tile_id // tl.cdiv(N, BLOCK_N) + pid_n = tile_id % tl.cdiv(N, BLOCK_N) + return pid_m, pid_n + +@triton.jit() +def mac_loop(A, B, C, + M, N, K, + locks, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, + iters_per_tile, + start_iter, end_iter, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + ACC_TYPE: tl.constexpr, GROUP_M: tl.constexpr): + + # where are we in the grid + tile_id = start_iter // iters_per_tile + if GROUP_M > 0: + pid_m, pid_n = swizzle_tile(tile_id, M, N, K, BLOCK_M, BLOCK_N, BLOCK_K, GROUP_M) + else: + pid_m, pid_n = linear_tile(tile_id, M, N, K, BLOCK_M, BLOCK_N, BLOCK_K, GROUP_M) + + rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + rk = tl.arange(0, BLOCK_K) + A = A + (rm[:, None] * stride_am + rk[None, :] * stride_ak) + BLOCK_K * stride_ak * (start_iter % iters_per_tile) + B = B + (rk[:, None] * stride_bk + rn[None, :] * stride_bn) + BLOCK_K * stride_bk * (start_iter % iters_per_tile) + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_TYPE) + + for current_iter in range(start_iter, end_iter): + a = tl.load(A) + b = tl.load(B) + acc += tl.dot(a, b) + A += BLOCK_K * stride_ak + B += BLOCK_K * stride_bk + acc = acc.to(C.dtype.element_ty) # restore C.dtype.element_ty + if end_iter % iters_per_tile == 0: # last iteration of the tile always happens before its start on another SM + C_ = C + (rm[:, None] * stride_cm + rn[None, :] * stride_cn) # compute inside the if/else to avoid spilling! + tl.store(C_, acc) + if start_iter % iters_per_tile != 0: # only if tile has been partially processed + tl.atomic_xchg(locks + tile_id, 1) + else: + while tl.atomic_cas(locks + tile_id, 1, 1) != 1: + pass + C_ = C + (rm[:, None] * stride_cm + rn[None, :] * stride_cn) # compute inside the if/else to avoid spilling! + tl.atomic_add(C_, acc) + +@triton.jit() +def first_wave( + A, B, C, + M, N, K, + locks, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, + total_full_tiles_streamk, total_partial_tiles_streamk, iters_per_tile, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ACC_TYPE: tl.constexpr, + GROUP_M: tl.constexpr, +): + pid = tl.program_id(0) + start_iter = pid * total_full_tiles_streamk + tl.minimum(pid, total_partial_tiles_streamk) + last_iter = (pid + 1) * total_full_tiles_streamk + tl.minimum(pid + 1, total_partial_tiles_streamk) + + while start_iter < last_iter: + end_iter = tl.minimum(start_iter + (iters_per_tile - start_iter % iters_per_tile), last_iter) + mac_loop(A, B, C, + M, N, K, + locks, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, + iters_per_tile, + start_iter, end_iter, + BLOCK_M, BLOCK_N, BLOCK_K, ACC_TYPE, + GROUP_M, + ) + + start_iter = end_iter + +@triton.jit() +def full_tiles( + A, B, C, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, + total_tiles_streamk, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ACC_TYPE: tl.constexpr, + GROUP_M: tl.constexpr, +): + # first wave has done more tiles than there are SMs, we adjust pid + tile_id = tl.program_id(0) + total_tiles_streamk + if GROUP_M > 0: + pid_m, pid_n = swizzle_tile(tile_id, M, N, K, BLOCK_M, BLOCK_N, BLOCK_K, GROUP_M) + else: + pid_m, pid_n = linear_tile(tile_id, M, N, K, BLOCK_M, BLOCK_N, BLOCK_K, GROUP_M) + + # do matrix multiplication + rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + rk = tl.arange(0, BLOCK_K) + # pointers + A = A + (rm[:, None] * stride_am + rk[None, :] * stride_ak) + B = B + (rk[:, None] * stride_bk + rn[None, :] * stride_bn) + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_TYPE) + for k in range(0, tl.cdiv(K, BLOCK_K)): + a = tl.load(A) + b = tl.load(B) + acc += tl.dot(a, b) + A += BLOCK_K * stride_ak + B += BLOCK_K * stride_bk + acc = acc.to(C.dtype.element_ty) # restore C.dtype.element_ty + # rematerialize rm and rn to save registers + rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + C = C + (rm[:, None] * stride_cm + rn[None, :] * stride_cn) + tl.store(C, acc) + +def run(a: torch.Tensor, b: torch.Tensor, block_size: int = None, + grid_programs: int = 108, # Default to roughly standard SM count, adjust via config + BLK_M: int = 128, BLK_N: int = 128, BLK_K: int = 32, + two_tiles: bool = True, num_stages: int = 3, num_warps: int = 4): + """ + Wrapper function to launch the Triton Stream-K kernels. + """ + device = a.device + + assert a.is_contiguous() and b.is_contiguous(), "non-contiguous inputs are not supported" + assert a.shape[1] == b.shape[0], "incompatible dimensions" + + M, K = a.shape + _, N = b.shape + + # accumulator types + ACC_TYPE = tl.float32 if a.dtype in[torch.float16, torch.bfloat16, torch.float32] else tl.int32 + + # compute grid configurations + total_blocks_M = triton.cdiv(M, BLK_M) + total_blocks_N = triton.cdiv(N, BLK_N) + iters_per_tile = triton.cdiv(K, BLK_K) + GROUP_M = 8 + total_tiles = total_blocks_M * total_blocks_N + + total_programs_streamk = grid_programs + + if total_programs_streamk > 0: + total_tiles_streamk = total_tiles % total_programs_streamk + if two_tiles and total_tiles - total_tiles_streamk > total_programs_streamk: + total_tiles_streamk += total_programs_streamk + + total_blocking_tiles = total_tiles - total_tiles_streamk + total_iters_streamk = total_tiles_streamk * iters_per_tile + total_full_tiles_streamk = total_iters_streamk // total_programs_streamk + total_partial_tiles_streamk = total_iters_streamk % total_programs_streamk + else: + total_blocking_tiles = total_tiles + total_tiles_streamk = 0 + total_full_tiles_streamk = 0 + total_partial_tiles_streamk = 0 + total_iters_streamk = 0 + + c = torch.empty((M, N), device=device, dtype=a.dtype) + locks = torch.zeros((total_tiles_streamk,), device=device, dtype=torch.int32) + + if total_programs_streamk > 0: + first_wave[(total_programs_streamk,)]( + a, b, c, M, N, K, locks, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), c.stride(1), + total_full_tiles_streamk=total_full_tiles_streamk, + total_partial_tiles_streamk=total_partial_tiles_streamk, + iters_per_tile=iters_per_tile, + BLOCK_M=BLK_M, BLOCK_N=BLK_N, BLOCK_K=BLK_K, + ACC_TYPE=ACC_TYPE, GROUP_M=GROUP_M, + num_stages=num_stages, num_warps=num_warps, + ) + + if total_blocking_tiles > 0: + full_tiles[(total_blocking_tiles,)]( + a, b, c, M, N, K, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), c.stride(1), + total_tiles_streamk=total_tiles_streamk, + BLOCK_M=BLK_M, BLOCK_N=BLK_N, BLOCK_K=BLK_K, + ACC_TYPE=ACC_TYPE, GROUP_M=GROUP_M, + num_stages=num_stages, num_warps=num_warps, + ) + + return c \ No newline at end of file diff --git a/core/verifier.py b/core/verifier.py index 11b807c3..9e81fda0 100644 --- a/core/verifier.py +++ b/core/verifier.py @@ -1,6 +1,6 @@ import torch -def verify(output, reference, atol=1e-3, rtol=1e-3): +def verify(output, reference, atol=2e-3, rtol=2e-3): try: torch.testing.assert_close(output, reference, atol=atol, rtol=rtol) return True, "" diff --git a/data/tensors.py b/data/tensors.py index dcece398..ac1b090a 100644 --- a/data/tensors.py +++ b/data/tensors.py @@ -1,4 +1,5 @@ import torch +import math def generate_vector_add_inputs(n, dtype=torch.float32, device='cuda'): x = torch.randn(n, dtype=dtype, device=device) @@ -9,6 +10,125 @@ def generate_vector_add_inputs(n, dtype=torch.float32, device='cuda'): def generate_sin_inputs(n, dtype=torch.float32, device='cuda'): x = torch.randn(n, dtype=dtype, device=device) return (x,) +def generate_rope_inputs(batch_size, seq_len, n_heads, head_dim, dtype=torch.float32, device='cuda', **kwargs): + + q = torch.randn(batch_size, seq_len, n_heads, head_dim, dtype=dtype, device=device) + + half_dim = head_dim // 2 + cos = torch.randn(seq_len, half_dim, dtype=dtype, device=device) + sin = torch.randn(seq_len, half_dim, dtype=dtype, device=device) + + return (q, cos, sin) +def generate_softmax_inputs(n=None, shape=None, dtype=torch.float32, device='cuda'): + if shape is None: + if n is None: + raise ValueError("Must provide 'n' or 'shape' for softmax inputs") + + cols = int(n**0.5) + rows = n // cols + shape = (rows, cols) + + if isinstance(dtype, str): + dtype = getattr(torch, dtype) + + x = torch.randn(*shape, dtype=dtype, device=device) + + + return (x,) +def generate_flash_attn_inputs(batch_size, n_heads, seq_len, head_dim, dtype=torch.float16, device='cuda', **kwargs): + + q = torch.randn(batch_size, n_heads, seq_len, head_dim, dtype=dtype, device=device) + k = torch.randn(batch_size, n_heads, seq_len, head_dim, dtype=dtype, device=device) + v = torch.randn(batch_size, n_heads, seq_len, head_dim, dtype=dtype, device=device) + + return (q.contiguous(), k.contiguous(), v.contiguous()) +def generate_flash_decode_stage2_inputs(n=None, batch=2, heads=8, seq_len=4096, head_dim=128, block_seq=128, dtype=torch.float32, device='cuda', **kwargs): + num_blocks = (seq_len + block_seq - 1) // block_seq + + b_seqlen = torch.full((batch,), seq_len, dtype=torch.int32, device=device) + mid_o = torch.randn((batch, heads, num_blocks, head_dim), dtype=dtype, device=device) + mid_o_lse = torch.randn((batch, heads, num_blocks), dtype=dtype, device=device) + + block_seq_tensor = torch.tensor(block_seq, dtype=torch.int32, device='cpu') + + return (mid_o, mid_o_lse, b_seqlen, block_seq_tensor) +def generate_mat_mul_inputs(M=1024, K=1024, N=1024, dtype=torch.float16, device='cuda', **kwargs): + # Depending on the PyTorch version, directly initializing float8_e4m3fn with randn might not be supported. + # The safest way is to generate float32 and cast. + a = torch.randn((M, K), dtype=torch.float32, device=device).to(dtype) + b = torch.randn((K, N), dtype=torch.float32, device=device).to(dtype) + return (a, b) + +def generate_mat_mul_int8_inputs(n=None, M=1024, N=1024, K_b=256, device='cuda', **kwargs): + K = K_b * 4 + + a = torch.randint(-128, 127, (M, K), dtype=torch.int8, device=device) + b = torch.randint(0, 255, (K_b, N), dtype=torch.uint8, device=device).to(torch.int8) + + return (a, b) + +def generate_streamk_matmul_inputs(n=None, M=1024, N=1024, K=1024, dtype=torch.float16, device='cuda', **kwargs): + if isinstance(dtype, str): + dtype = getattr(torch, dtype) + + a = torch.randn((M, K), dtype=torch.float32, device=device).to(dtype) + b = torch.randn((K, N), dtype=torch.float32, device=device).to(dtype) + + return (a, b) +def generate_block_sparse_attention_inputs(n=None, 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=torch.float16, device='cuda', **kwargs): + """ + Generate inputs for block sparse attention. + Creates a simple "Local Window + Causal" sparse CSR layout. + """ + if isinstance(dtype, str): + dtype = getattr(torch, dtype) + + Q = torch.randn((B, H, M, D), dtype=dtype, device=device) + K = torch.randn((B, H_kv, M, D), dtype=dtype, device=device) # N == M + V = torch.randn((B, H_kv, M, D), dtype=dtype, device=device) + + num_layout = 1 # Shared layout for all heads + num_rows = math.ceil(M / BLOCK_M) + num_cols = math.ceil(M / BLOCK_N) + + layout_csr_row_stride_h = num_rows + 1 + layout_csr_col_stride_h = num_rows * num_cols # Max possible capacity + + # We build a causal local window mask + window_blocks = 2 # Attend to current block and 2 previous blocks + + row_ptrs = [] + col_indices =[] + + current_ptr = 0 + for r in range(num_rows): + row_ptrs.append(current_ptr) + # Start col is max(0, r - window_blocks) + # End col is r (inclusive, because of causal) + start_c = max(0, r - window_blocks) + end_c = r + for c in range(start_c, end_c + 1): + col_indices.append(c) + current_ptr += 1 + + row_ptrs.append(current_ptr) # Final ptr + + # Pad col_indices to required size + col_indices = col_indices + [0] * (layout_csr_col_stride_h - len(col_indices)) + + layout_csr_row_indices = torch.tensor(row_ptrs, dtype=torch.int32, device=device) + layout_csr_col_indices = torch.tensor(col_indices, dtype=torch.int32, device=device) + + softmax_scale = 1.0 / math.sqrt(D) + EVEN_M = (M % BLOCK_M == 0) + EVEN_N = (M % BLOCK_N == 0) + + return (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, H, H_kv, M, + BLOCK_M, EVEN_M, BLOCK_N, EVEN_N, BLOCK_D, NUM_D_BLOCKS) def generate_swiglu_inputs(batch_size, ncols, dtype=torch.float32, device='cuda'): x = torch.randn(batch_size, ncols, dtype=dtype, device=device) @@ -74,6 +194,14 @@ def generate_generic_fused_container_inputs(n, dtype=torch.float32, device='cuda GENERATORS = { "vector_add": generate_vector_add_inputs, "sin": generate_sin_inputs, + "rope": generate_rope_inputs, + "flash_attention": generate_flash_attn_inputs, + "softmax": generate_softmax_inputs, + "flash_decode": generate_flash_decode_stage2_inputs, + "matmul_fp16_fp8": generate_mat_mul_inputs, + "matmul_int8": generate_mat_mul_int8_inputs, + "streamk_matmul": generate_streamk_matmul_inputs, + "block_sparse_attention": generate_block_sparse_attention_inputs, "cross_entropy": generate_cross_entropy_inputs, "matrix_transpose": generate_matrix_transpose_inputs, "swiglu": generate_swiglu_inputs,