gpt-5 / triton7308c5
gpt-5_triton_7308c5 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 386 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-7308c5?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
21 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #cb4b00
NVIDIA B200
146.3µs
#5 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
147.3µs
#8 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #5fa6fa
NVIDIA B200
147.4µs
#5 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
148.1µs
#9 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
148.3µs
#9 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #f3d59b
NVIDIA B200
148.8µs
#3 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
149.0µs
#9 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
149.9µs
#10 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
150.0µs
#10 of 20
2025-10-19
Show all 21 measurements ›Showing all 21 measurements ⌄
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
150.2µs
#11 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
150.5µs
#12 of 20
2025-10-19
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:561ea20168ac707b05662f1c596a5f16c3ebf9160e26614c28d99fe0601509a5
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
main.py386 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def gqa_ragged_prefill_causal_h32_kv8_d128_kernel(
q_ptr, k_ptr, v_ptr,
stride_q_q, stride_q_h, stride_q_d,
stride_k_k, stride_k_h, stride_k_d,
stride_v_k, stride_v_h, stride_v_d,
out_ptr, stride_out_q, stride_out_h, stride_out_d,
lse_ptr, stride_lse_q, stride_lse_h,
q_kv_start_ptr, q_kv_max_ptr,
total_q,
sm_scale, ln2,
RATIO: tl.constexpr, HEAD_DIM: tl.constexpr,
BLOCK_N: tl.constexpr, BLOCK_DK: tl.constexpr, BLOCK_DV: tl.constexpr,
):
pid_q = tl.program_id(0)
kvh = tl.program_id(1)
if pid_q >= total_q:
return
kv_start = tl.load(q_kv_start_ptr + pid_q, mask=True, other=0).to(tl.int32)
kv_max = tl.load(q_kv_max_ptr + pid_q, mask=True, other=0).to(tl.int32)
heads_base = kvh * RATIO
neg_inf = tl.full([], -float("inf"), tl.float32)
if kv_max <= 0:
# No available keys for this query; set LSE to -inf and outputs to 0
for r in range(RATIO):
lse_ptr_r = lse_ptr + pid_q * stride_lse_q + (heads_base + r) * stride_lse_h
tl.store(lse_ptr_r, neg_inf)
# store output zeros
for dv0 in range(0, HEAD_DIM, BLOCK_DV):
d_voffs = dv0 + tl.arange(0, BLOCK_DV)
out_ptrs = out_ptr + pid_q * stride_out_q + (heads_base + r) * stride_out_h + d_voffs * stride_out_d
tl.store(out_ptrs, tl.zeros([BLOCK_DV], dtype=tl.bfloat16))
return
# Initialize streaming softmax stats per head (RATIO=4)
m0 = neg_inf
m1 = neg_inf
m2 = neg_inf
m3 = neg_inf
l0 = tl.zeros([], dtype=tl.float32)
l1 = tl.zeros([], dtype=tl.float32)
l2 = tl.zeros([], dtype=tl.float32)
l3 = tl.zeros([], dtype=tl.float32)
# Output accumulators per head, split along D into 4 segments (HEAD_DIM=128, BLOCK_DV=32)
o0_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o0_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o0_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o0_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o1_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o1_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o1_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o1_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o2_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o2_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o2_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o2_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o3_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o3_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o3_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
o3_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)
# Loop over key tiles
for start_n in range(0, kv_max, BLOCK_N):
key_offsets = start_n + tl.arange(0, BLOCK_N)
key_mask = key_offsets < kv_max
# Accumulate logits per head for this tile
logits0 = tl.zeros([BLOCK_N], dtype=tl.float32)
logits1 = tl.zeros([BLOCK_N], dtype=tl.float32)
logits2 = tl.zeros([BLOCK_N], dtype=tl.float32)
logits3 = tl.zeros([BLOCK_N], dtype=tl.float32)
for d0 in range(0, HEAD_DIM, BLOCK_DK):
d_off = d0 + tl.arange(0, BLOCK_DK)
# Load K chunk: [BLOCK_N, BLOCK_DK] -> fp32
k_ptrs = k_ptr + (kv_start + key_offsets)[:, None] * stride_k_k + kvh * stride_k_h + d_off[None, :] * stride_k_d
k_chunk = tl.load(
k_ptrs,
mask=key_mask[:, None] & (d_off[None, :] < HEAD_DIM),
other=0
).to(tl.float32)
# Load Q chunk and accumulate logits for each of the 4 query heads in this kv head group
# Head 0
q_ptrs0 = q_ptr + pid_q * stride_q_q + (heads_base + 0) * stride_q_h + d_off * stride_q_d
q_vec0 = tl.load(q_ptrs0, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
logits0 += tl.sum(k_chunk * q_vec0[None, :], axis=1)
# Head 1
q_ptrs1 = q_ptr + pid_q * stride_q_q + (heads_base + 1) * stride_q_h + d_off * stride_q_d
q_vec1 = tl.load(q_ptrs1, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
logits1 += tl.sum(k_chunk * q_vec1[None, :], axis=1)
# Head 2
q_ptrs2 = q_ptr + pid_q * stride_q_q + (heads_base + 2) * stride_q_h + d_off * stride_q_d
q_vec2 = tl.load(q_ptrs2, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
logits2 += tl.sum(k_chunk * q_vec2[None, :], axis=1)
# Head 3
q_ptrs3 = q_ptr + pid_q * stride_q_q + (heads_base + 3) * stride_q_h + d_off * stride_q_d
q_vec3 = tl.load(q_ptrs3, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
logits3 += tl.sum(k_chunk * q_vec3[None, :], axis=1)
# Scale and apply mask
p0 = logits0 * sm_scale
p1 = logits1 * sm_scale
p2 = logits2 * sm_scale
p3 = logits3 * sm_scale
p0 = tl.where(key_mask, p0, neg_inf)
p1 = tl.where(key_mask, p1, neg_inf)
p2 = tl.where(key_mask, p2, neg_inf)
p3 = tl.where(key_mask, p3, neg_inf)
# Preload V chunks once for the tile (reused across heads)
d0_idx = 0 + tl.arange(0, BLOCK_DV)
d1_idx = BLOCK_DV + tl.arange(0, BLOCK_DV)
d2_idx = 2 * BLOCK_DV + tl.arange(0, BLOCK_DV)
d3_idx = 3 * BLOCK_DV + tl.arange(0, BLOCK_DV)
v_ptrs0 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d0_idx[None, :] * stride_v_d
v_ptrs1 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d1_idx[None, :] * stride_v_d
v_ptrs2 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d2_idx[None, :] * stride_v_d
v_ptrs3 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d3_idx[None, :] * stride_v_d
v_chunk0 = tl.load(v_ptrs0, mask=key_mask[:, None], other=0).to(tl.float32)
v_chunk1 = tl.load(v_ptrs1, mask=key_mask[:, None], other=0).to(tl.float32)
v_chunk2 = tl.load(v_ptrs2, mask=key_mask[:, None], other=0).to(tl.float32)
v_chunk3 = tl.load(v_ptrs3, mask=key_mask[:, None], other=0).to(tl.float32)
# Head 0
m0_tile = tl.max(p0, axis=0)
m0_new = tl.maximum(m0, m0_tile)
alpha0 = tl.exp(m0 - m0_new)
o0_s0 = o0_s0 * alpha0
o0_s1 = o0_s1 * alpha0
o0_s2 = o0_s2 * alpha0
o0_s3 = o0_s3 * alpha0
w0 = tl.exp(p0 - m0_new)
l0 = l0 * alpha0 + tl.sum(w0, axis=0)
o0_s0 = o0_s0 + tl.sum(v_chunk0 * w0[:, None], axis=0)
o0_s1 = o0_s1 + tl.sum(v_chunk1 * w0[:, None], axis=0)
o0_s2 = o0_s2 + tl.sum(v_chunk2 * w0[:, None], axis=0)
o0_s3 = o0_s3 + tl.sum(v_chunk3 * w0[:, None], axis=0)
m0 = m0_new
# Head 1
m1_tile = tl.max(p1, axis=0)
m1_new = tl.maximum(m1, m1_tile)
alpha1 = tl.exp(m1 - m1_new)
o1_s0 = o1_s0 * alpha1
o1_s1 = o1_s1 * alpha1
o1_s2 = o1_s2 * alpha1
o1_s3 = o1_s3 * alpha1
w1 = tl.exp(p1 - m1_new)
l1 = l1 * alpha1 + tl.sum(w1, axis=0)
o1_s0 = o1_s0 + tl.sum(v_chunk0 * w1[:, None], axis=0)
o1_s1 = o1_s1 + tl.sum(v_chunk1 * w1[:, None], axis=0)
o1_s2 = o1_s2 + tl.sum(v_chunk2 * w1[:, None], axis=0)
o1_s3 = o1_s3 + tl.sum(v_chunk3 * w1[:, None], axis=0)
m1 = m1_new
# Head 2
m2_tile = tl.max(p2, axis=0)
m2_new = tl.maximum(m2, m2_tile)
alpha2 = tl.exp(m2 - m2_new)
o2_s0 = o2_s0 * alpha2
o2_s1 = o2_s1 * alpha2
o2_s2 = o2_s2 * alpha2
o2_s3 = o2_s3 * alpha2
w2 = tl.exp(p2 - m2_new)
l2 = l2 * alpha2 + tl.sum(w2, axis=0)
o2_s0 = o2_s0 + tl.sum(v_chunk0 * w2[:, None], axis=0)
o2_s1 = o2_s1 + tl.sum(v_chunk1 * w2[:, None], axis=0)
o2_s2 = o2_s2 + tl.sum(v_chunk2 * w2[:, None], axis=0)
o2_s3 = o2_s3 + tl.sum(v_chunk3 * w2[:, None], axis=0)
m2 = m2_new
# Head 3
m3_tile = tl.max(p3, axis=0)
m3_new = tl.maximum(m3, m3_tile)
alpha3 = tl.exp(m3 - m3_new)
o3_s0 = o3_s0 * alpha3
o3_s1 = o3_s1 * alpha3
o3_s2 = o3_s2 * alpha3
o3_s3 = o3_s3 * alpha3
w3 = tl.exp(p3 - m3_new)
l3 = l3 * alpha3 + tl.sum(w3, axis=0)
o3_s0 = o3_s0 + tl.sum(v_chunk0 * w3[:, None], axis=0)
o3_s1 = o3_s1 + tl.sum(v_chunk1 * w3[:, None], axis=0)
o3_s2 = o3_s2 + tl.sum(v_chunk2 * w3[:, None], axis=0)
o3_s3 = o3_s3 + tl.sum(v_chunk3 * w3[:, None], axis=0)
m3 = m3_new
# Finalize: compute output = O / l, and lse = (m + log(l)) / ln2
d0 = 0 + tl.arange(0, BLOCK_DV)
d1 = BLOCK_DV + tl.arange(0, BLOCK_DV)
d2 = 2 * BLOCK_DV + tl.arange(0, BLOCK_DV)
d3 = 3 * BLOCK_DV + tl.arange(0, BLOCK_DV)
# Head 0
l0_pos = l0 > 0
lse0 = tl.where(l0_pos, (m0 + tl.log(l0)) / ln2, neg_inf)
lse_ptr0 = lse_ptr + pid_q * stride_lse_q + (heads_base + 0) * stride_lse_h
tl.store(lse_ptr0, lse0)
out_ptrs0 = out_ptr + pid_q * stride_out_q + (heads_base + 0) * stride_out_h
o0_s0_out = tl.where(l0_pos, o0_s0 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
o0_s1_out = tl.where(l0_pos, o0_s1 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
o0_s2_out = tl.where(l0_pos, o0_s2 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
o0_s3_out = tl.where(l0_pos, o0_s3 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
tl.store(out_ptrs0 + d0 * stride_out_d, o0_s0_out.to(tl.bfloat16))
tl.store(out_ptrs0 + d1 * stride_out_d, o0_s1_out.to(tl.bfloat16))
tl.store(out_ptrs0 + d2 * stride_out_d, o0_s2_out.to(tl.bfloat16))
tl.store(out_ptrs0 + d3 * stride_out_d, o0_s3_out.to(tl.bfloat16))
# Head 1
l1_pos = l1 > 0
lse1 = tl.where(l1_pos, (m1 + tl.log(l1)) / ln2, neg_inf)
lse_ptr1 = lse_ptr + pid_q * stride_lse_q + (heads_base + 1) * stride_lse_h
tl.store(lse_ptr1, lse1)
out_ptrs1 = out_ptr + pid_q * stride_out_q + (heads_base + 1) * stride_out_h
o1_s0_out = tl.where(l1_pos, o1_s0 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
o1_s1_out = tl.where(l1_pos, o1_s1 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
o1_s2_out = tl.where(l1_pos, o1_s2 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
o1_s3_out = tl.where(l1_pos, o1_s3 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
tl.store(out_ptrs1 + d0 * stride_out_d, o1_s0_out.to(tl.bfloat16))
tl.store(out_ptrs1 + d1 * stride_out_d, o1_s1_out.to(tl.bfloat16))
tl.store(out_ptrs1 + d2 * stride_out_d, o1_s2_out.to(tl.bfloat16))
tl.store(out_ptrs1 + d3 * stride_out_d, o1_s3_out.to(tl.bfloat16))
# Head 2
l2_pos = l2 > 0
lse2 = tl.where(l2_pos, (m2 + tl.log(l2)) / ln2, neg_inf)
lse_ptr2 = lse_ptr + pid_q * stride_lse_q + (heads_base + 2) * stride_lse_h
tl.store(lse_ptr2, lse2)
out_ptrs2 = out_ptr + pid_q * stride_out_q + (heads_base + 2) * stride_out_h
o2_s0_out = tl.where(l2_pos, o2_s0 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
o2_s1_out = tl.where(l2_pos, o2_s1 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
o2_s2_out = tl.where(l2_pos, o2_s2 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
o2_s3_out = tl.where(l2_pos, o2_s3 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
tl.store(out_ptrs2 + d0 * stride_out_d, o2_s0_out.to(tl.bfloat16))
tl.store(out_ptrs2 + d1 * stride_out_d, o2_s1_out.to(tl.bfloat16))
tl.store(out_ptrs2 + d2 * stride_out_d, o2_s2_out.to(tl.bfloat16))
tl.store(out_ptrs2 + d3 * stride_out_d, o2_s3_out.to(tl.bfloat16))
# Head 3
l3_pos = l3 > 0
lse3 = tl.where(l3_pos, (m3 + tl.log(l3)) / ln2, neg_inf)
lse_ptr3 = lse_ptr + pid_q * stride_lse_q + (heads_base + 3) * stride_lse_h
tl.store(lse_ptr3, lse3)
out_ptrs3 = out_ptr + pid_q * stride_out_q + (heads_base + 3) * stride_out_h
o3_s0_out = tl.where(l3_pos, o3_s0 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
o3_s1_out = tl.where(l3_pos, o3_s1 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
o3_s2_out = tl.where(l3_pos, o3_s2 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
o3_s3_out = tl.where(l3_pos, o3_s3 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
tl.store(out_ptrs3 + d0 * stride_out_d, o3_s0_out.to(tl.bfloat16))
tl.store(out_ptrs3 + d1 * stride_out_d, o3_s1_out.to(tl.bfloat16))
tl.store(out_ptrs3 + d2 * stride_out_d, o3_s2_out.to(tl.bfloat16))
tl.store(out_ptrs3 + d3 * stride_out_d, o3_s3_out.to(tl.bfloat16))
def _prepare_q_meta_from_indptr(qo_indptr: torch.Tensor, kv_indptr: torch.Tensor):
# Build per-query arrays: kv_start[q], kv_max[q]
qo_indptr_cpu = qo_indptr.to("cpu", non_blocking=False)
kv_indptr_cpu = kv_indptr.to("cpu", non_blocking=False)
len_indptr = qo_indptr_cpu.numel()
total_q = int(qo_indptr_cpu[-1].item())
q_kv_start = torch.empty(total_q, dtype=torch.int32)
q_kv_max = torch.empty(total_q, dtype=torch.int32)
for b in range(len_indptr - 1):
q_start = int(qo_indptr_cpu[b].item())
q_end = int(qo_indptr_cpu[b + 1].item())
kv_start = int(kv_indptr_cpu[b].item())
kv_end = int(kv_indptr_cpu[b + 1].item())
q_len = q_end - q_start
kv_len = kv_end - kv_start
if q_len <= 0:
continue
delta = kv_len - q_len
pos = torch.arange(q_len, dtype=torch.int32)
kv_max = pos + 1 + int(delta)
kv_max = torch.clamp(kv_max, min=0, max=kv_len)
q_kv_start[q_start:q_end] = int(kv_start)
q_kv_max[q_start:q_end] = kv_max
return q_kv_start, q_kv_max
@torch.no_grad()
def run(q, k, v, qo_indptr, kv_indptr, sm_scale=None):
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required to run Triton kernels. No CUDA device is available.")
HEAD_DIM = 128
NUM_QO = 32
NUM_KV = 8
RATIO = NUM_QO // NUM_KV # 4
inputs = [q, k, v, qo_indptr, kv_indptr]
orig_devices = [t.device for t in inputs]
target_device = None
for t in inputs:
if t.is_cuda:
target_device = t.device
break
if target_device is None:
target_device = torch.device("cuda")
q = q.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
k = k.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
v = v.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
qo_indptr = qo_indptr.to(device=target_device, dtype=torch.int32, non_blocking=True)
kv_indptr = kv_indptr.to(device=target_device, dtype=torch.int32, non_blocking=True)
total_q, num_qo_heads, head_dim = q.shape
total_kv, num_kv_heads, _ = k.shape
if num_qo_heads != NUM_QO:
raise ValueError(f"num_qo_heads must be {NUM_QO}, got {num_qo_heads}")
if num_kv_heads != NUM_KV:
raise ValueError(f"num_kv_heads must be {NUM_KV}, got {num_kv_heads}")
if head_dim != HEAD_DIM:
raise ValueError(f"head_dim must be {HEAD_DIM}, got {head_dim}")
if int(qo_indptr[-1].item()) != total_q:
raise ValueError("Constraint violated: total_q must equal qo_indptr[-1]")
if int(kv_indptr[-1].item()) != total_kv:
raise ValueError("Constraint violated: total_kv must equal kv_indptr[-1]")
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(HEAD_DIM)
sm_scale = float(sm_scale)
ln2 = float(math.log(2.0))
# Prepare per-query kv_start and kv_max on CPU for simplicity, then move to target device
qo_indptr_cpu = qo_indptr.to("cpu")
kv_indptr_cpu = kv_indptr.to("cpu")
q_kv_start_cpu, q_kv_max_cpu = _prepare_q_meta_from_indptr(qo_indptr_cpu, kv_indptr_cpu)
q_kv_start = q_kv_start_cpu.to(device=target_device, non_blocking=True)
q_kv_max = q_kv_max_cpu.to(device=target_device, non_blocking=True)
out_gpu = torch.empty((total_q, NUM_QO, HEAD_DIM), dtype=torch.bfloat16, device=target_device)
lse_gpu = torch.empty((total_q, NUM_QO), dtype=torch.float32, device=target_device)
stride_q_q, stride_q_h, stride_q_d = q.stride()
stride_k_k, stride_k_h, stride_k_d = k.stride()
stride_v_k, stride_v_h, stride_v_d = v.stride()
stride_out_q, stride_out_h, stride_out_d = out_gpu.stride()
stride_lse_q, stride_lse_h = lse_gpu.stride()
grid = (total_q, NUM_KV)
BLOCK_N = 64
BLOCK_DK = 32
BLOCK_DV = 32
gqa_ragged_prefill_causal_h32_kv8_d128_kernel[grid](
q, k, v,
stride_q_q, stride_q_h, stride_q_d,
stride_k_k, stride_k_h, stride_k_d,
stride_v_k, stride_v_h, stride_v_d,
out_gpu, stride_out_q, stride_out_h, stride_out_d,
lse_gpu, stride_lse_q, stride_lse_h,
q_kv_start, q_kv_max,
total_q,
sm_scale, ln2,
RATIO=RATIO, HEAD_DIM=HEAD_DIM,
BLOCK_N=BLOCK_N, BLOCK_DK=BLOCK_DK, BLOCK_DV=BLOCK_DV,
num_warps=4, num_stages=2,
)
out = out_gpu.to(orig_devices[0], non_blocking=True)
lse = lse_gpu.to(orig_devices[0], non_blocking=True)
return out, lsescrolls · 386 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON