gpt-5 / tritoncb1275
gpt-5_triton_cb1275 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 344 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-cb1275?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
48 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
86.8µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
87.4µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=55 · num_kv_indices=38
NVIDIA B200
87.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
88.6µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
88.7µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
89.3µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
89.7µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=63 · num_kv_indices=46
NVIDIA B200
92.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
92.7µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=11 · num_kv_indices=10
NVIDIA B200
93.8µs
#7 of 7
2025-10-16
Show all 48 measurements ›Showing all 48 measurements ⌄
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
94.4µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
94.4µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
94.9µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
102.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
112.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
113.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
113.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1396 · num_kv_indices=1369
NVIDIA B200
114.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
114.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
114.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
114.4µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1732 · num_kv_indices=1705
NVIDIA B200
115.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
116.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2212 · num_kv_indices=2185
NVIDIA B200
117.8µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
122.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
123.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
123.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
124.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2708 · num_kv_indices=2681
NVIDIA B200
125.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=4333 · num_kv_indices=4298
NVIDIA B200
133.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
212.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
250.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
435.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
437.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
441.4µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
444.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65037 · num_kv_indices=55766
NVIDIA B200
444.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
447.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64205 · num_kv_indices=54934
NVIDIA B200
447.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
448.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67021 · num_kv_indices=57750
NVIDIA B200
448.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
450.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
452.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
452.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67405 · num_kv_indices=58134
NVIDIA B200
458.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
462.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=68237 · num_kv_indices=58966
NVIDIA B200
467.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63437 · num_kv_indices=54166
NVIDIA B200
469.2µs
#5 of 7
2025-10-16
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:10231abe2f3401b0219cd1d797d674a685a6d3a88856dbfa7eb880b115a36a37
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.
num-warps = 8
num_warps=8, num_stages=3,stages = 3
num_warps=8, num_stages=3,Kernel source
main.py344 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def gqa_paged_decode_h32_kv8_d128_ps1_kernel(
q_ptr, k_ptr, v_ptr,
kv_indptr_ptr, kv_indices_ptr,
out_ptr, lse_ptr,
sm_scale_ptr, inv_ln2_ptr,
batch_size,
stride_q_b, stride_q_h, stride_q_d,
stride_k_p, stride_k_ps, stride_k_h, stride_k_d,
stride_v_p, stride_v_ps, stride_v_h, stride_v_d,
stride_out_b, stride_out_h, stride_out_d,
stride_lse_b, stride_lse_h,
BLOCK_T: tl.constexpr, BLOCK_D: tl.constexpr, STEP: tl.constexpr,
GQA_RATIO: tl.constexpr,
):
pid = tl.program_id(0)
num_qo_heads = 32
b = pid // num_qo_heads
h = pid % num_qo_heads
if b >= batch_size:
return
# Load scalar parameters
sm_scale = tl.load(sm_scale_ptr)
inv_ln2 = tl.load(inv_ln2_ptr)
b_i64 = b.to(tl.int64)
h_i64 = h.to(tl.int64)
# GQA mapping
kv_head = (h // GQA_RATIO)
kv_head_i64 = kv_head.to(tl.int64)
# Load start/end pointers for this batch element
page_start = tl.load(kv_indptr_ptr + b_i64)
page_end = tl.load(kv_indptr_ptr + b_i64 + 1)
n_tokens = page_end - page_start
# Prepare output/lse pointers
d_all = tl.arange(0, BLOCK_D)
out_row_ptrs = out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + d_all.to(tl.int64) * stride_out_d
lse_ptr_ = lse_ptr + b_i64 * stride_lse_b + h_i64 * stride_lse_h
# If no tokens, write zeros and -inf LSE and return
if n_tokens <= 0:
zero_bf16 = tl.zeros([BLOCK_D], dtype=tl.bfloat16)
tl.store(out_row_ptrs, zero_bf16, mask=d_all < BLOCK_D)
tl.store(lse_ptr_, -float("inf"))
return
# Preload Q in four STEP chunks (bf16 -> fp32)
# chunk 0
d_idx0 = tl.arange(0, STEP)
q0 = tl.load(
q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx0.to(tl.int64) * stride_q_d,
mask=d_idx0 < BLOCK_D,
other=0,
).to(tl.float32)
# chunk 1
d_idx1 = STEP + tl.arange(0, STEP)
q1 = tl.load(
q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx1.to(tl.int64) * stride_q_d,
mask=d_idx1 < BLOCK_D,
other=0,
).to(tl.float32)
# chunk 2
d_idx2 = (2 * STEP) + tl.arange(0, STEP)
q2 = tl.load(
q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx2.to(tl.int64) * stride_q_d,
mask=d_idx2 < BLOCK_D,
other=0,
).to(tl.float32)
# chunk 3
d_idx3 = (3 * STEP) + tl.arange(0, STEP)
q3 = tl.load(
q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx3.to(tl.int64) * stride_q_d,
mask=d_idx3 < BLOCK_D,
other=0,
).to(tl.float32)
# Streaming softmax variables
m = -float("inf")
l = 0.0
# Accumulator for output across head_dim in 4 chunks (STEP each)
out_acc0 = tl.zeros([STEP], dtype=tl.float32)
out_acc1 = tl.zeros([STEP], dtype=tl.float32)
out_acc2 = tl.zeros([STEP], dtype=tl.float32)
out_acc3 = tl.zeros([STEP], dtype=tl.float32)
# Loop over tokens in blocks of BLOCK_T
pos = 0
while pos < n_tokens:
t_offsets = tl.arange(0, BLOCK_T)
offs = pos + t_offsets
mask_t = offs < n_tokens
# Gather page indices for this block
idx = tl.load(kv_indices_ptr + page_start.to(tl.int64) + offs.to(tl.int64), mask=mask_t, other=0)
# Compute logits for this block: [BLOCK_T]
logits = tl.zeros([BLOCK_T], dtype=tl.float32)
# chunk 0
k_ptrs0 = (
k_ptr
+ (idx[:, None].to(tl.int64) * stride_k_p)
+ (kv_head_i64 * stride_k_h)
+ (d_idx0[None, :].to(tl.int64) * stride_k_d)
)
k0 = tl.load(k_ptrs0, mask=mask_t[:, None] & (d_idx0[None, :] < BLOCK_D), other=0).to(tl.float32)
logits += tl.sum(k0 * q0[None, :], axis=1)
# chunk 1
k_ptrs1 = (
k_ptr
+ (idx[:, None].to(tl.int64) * stride_k_p)
+ (kv_head_i64 * stride_k_h)
+ (d_idx1[None, :].to(tl.int64) * stride_k_d)
)
k1 = tl.load(k_ptrs1, mask=mask_t[:, None] & (d_idx1[None, :] < BLOCK_D), other=0).to(tl.float32)
logits += tl.sum(k1 * q1[None, :], axis=1)
# chunk 2
k_ptrs2 = (
k_ptr
+ (idx[:, None].to(tl.int64) * stride_k_p)
+ (kv_head_i64 * stride_k_h)
+ (d_idx2[None, :].to(tl.int64) * stride_k_d)
)
k2 = tl.load(k_ptrs2, mask=mask_t[:, None] & (d_idx2[None, :] < BLOCK_D), other=0).to(tl.float32)
logits += tl.sum(k2 * q2[None, :], axis=1)
# chunk 3
k_ptrs3 = (
k_ptr
+ (idx[:, None].to(tl.int64) * stride_k_p)
+ (kv_head_i64 * stride_k_h)
+ (d_idx3[None, :].to(tl.int64) * stride_k_d)
)
k3 = tl.load(k_ptrs3, mask=mask_t[:, None] & (d_idx3[None, :] < BLOCK_D), other=0).to(tl.float32)
logits += tl.sum(k3 * q3[None, :], axis=1)
# Scale logits and apply mask
logits = logits * sm_scale
logits = tl.where(mask_t, logits, -float("inf"))
# Compute block max and update running m and l
block_max = tl.max(logits, axis=0)
new_m = tl.maximum(m, block_max)
scale_old = tl.exp(m - new_m)
# Weights for this block
weights = tl.exp(logits - new_m)
# Update l
l = l * scale_old + tl.sum(weights, axis=0)
# Scale previous accumulators by scale_old
out_acc0 *= scale_old
out_acc1 *= scale_old
out_acc2 *= scale_old
out_acc3 *= scale_old
# Accumulate V weighted by weights
# chunk 0
v_ptrs0 = (
v_ptr
+ (idx[:, None].to(tl.int64) * stride_v_p)
+ (kv_head_i64 * stride_v_h)
+ (d_idx0[None, :].to(tl.int64) * stride_v_d)
)
v0 = tl.load(v_ptrs0, mask=mask_t[:, None] & (d_idx0[None, :] < BLOCK_D), other=0).to(tl.float32)
out_acc0 += tl.sum(v0 * weights[:, None], axis=0)
# chunk 1
v_ptrs1 = (
v_ptr
+ (idx[:, None].to(tl.int64) * stride_v_p)
+ (kv_head_i64 * stride_v_h)
+ (d_idx1[None, :].to(tl.int64) * stride_v_d)
)
v1 = tl.load(v_ptrs1, mask=mask_t[:, None] & (d_idx1[None, :] < BLOCK_D), other=0).to(tl.float32)
out_acc1 += tl.sum(v1 * weights[:, None], axis=0)
# chunk 2
v_ptrs2 = (
v_ptr
+ (idx[:, None].to(tl.int64) * stride_v_p)
+ (kv_head_i64 * stride_v_h)
+ (d_idx2[None, :].to(tl.int64) * stride_v_d)
)
v2 = tl.load(v_ptrs2, mask=mask_t[:, None] & (d_idx2[None, :] < BLOCK_D), other=0).to(tl.float32)
out_acc2 += tl.sum(v2 * weights[:, None], axis=0)
# chunk 3
v_ptrs3 = (
v_ptr
+ (idx[:, None].to(tl.int64) * stride_v_p)
+ (kv_head_i64 * stride_v_h)
+ (d_idx3[None, :].to(tl.int64) * stride_v_d)
)
v3 = tl.load(v_ptrs3, mask=mask_t[:, None] & (d_idx3[None, :] < BLOCK_D), other=0).to(tl.float32)
out_acc3 += tl.sum(v3 * weights[:, None], axis=0)
# Update running max
m = new_m
pos += BLOCK_T
# Finalize lse in base 2
lse_base2 = (tl.log(l) + m) * inv_ln2
tl.store(lse_ptr_, lse_base2)
# Normalize by l
inv_l = 1.0 / l
out_acc0 *= inv_l
out_acc1 *= inv_l
out_acc2 *= inv_l
out_acc3 *= inv_l
# Store output chunks
# chunk 0
tl.store(
out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + (tl.arange(0, STEP).to(tl.int64)) * stride_out_d,
out_acc0.to(tl.bfloat16),
mask=(tl.arange(0, STEP) < BLOCK_D),
)
# chunk 1
tl.store(
out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + ((STEP + tl.arange(0, STEP)).to(tl.int64)) * stride_out_d,
out_acc1.to(tl.bfloat16),
mask=((STEP + tl.arange(0, STEP)) < BLOCK_D),
)
# chunk 2
tl.store(
out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + (((2 * STEP) + tl.arange(0, STEP)).to(tl.int64)) * stride_out_d,
out_acc2.to(tl.bfloat16),
mask=(((2 * STEP) + tl.arange(0, STEP)) < BLOCK_D),
)
# chunk 3
tl.store(
out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + (((3 * STEP) + tl.arange(0, STEP)).to(tl.int64)) * stride_out_d,
out_acc3.to(tl.bfloat16),
mask=(((3 * STEP) + tl.arange(0, STEP)) < BLOCK_D),
)
def run(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale=None):
# Validate inputs and move to CUDA if available
if not torch.cuda.is_available():
# If any tensor is already on CUDA but CUDA is unavailable, raise error
if any(t.is_cuda for t in [q, k_cache, v_cache, kv_indptr, kv_indices] if isinstance(t, torch.Tensor)):
raise RuntimeError("CUDA is not available but some inputs are CUDA tensors.")
raise RuntimeError("CUDA is required to run Triton kernels. Please enable a CUDA-capable device.")
device_out = q.device
def to_cuda(t):
return t if t.is_cuda else t.cuda()
q_c = to_cuda(q)
k_c = to_cuda(k_cache)
v_c = to_cuda(v_cache)
kv_indptr_c = to_cuda(kv_indptr)
kv_indices_c = to_cuda(kv_indices)
# Check dtypes and shapes
assert q_c.dtype == torch.bfloat16, "q must be bfloat16"
assert k_c.dtype == torch.bfloat16 and v_c.dtype == torch.bfloat16, "k_cache and v_cache must be bfloat16"
assert kv_indptr_c.dtype == torch.int32, "kv_indptr must be int32"
assert kv_indices_c.dtype == torch.int32, "kv_indices must be int32"
batch_size, num_qo_heads, head_dim = q_c.shape
num_pages, page_size, num_kv_heads, head_dim_k = k_c.shape
assert num_qo_heads == 32, "num_qo_heads must be 32"
assert num_kv_heads == 8, "num_kv_heads must be 8"
assert head_dim == 128 and head_dim_k == 128, "head_dim must be 128"
assert page_size == 1, "page_size must be 1"
len_indptr = kv_indptr_c.shape[0]
num_kv_indices = kv_indices_c.shape[0]
assert len_indptr == batch_size + 1, "len_indptr must equal batch_size + 1"
last = kv_indptr_c[-1].item()
assert num_kv_indices == last, "num_kv_indices must equal kv_indptr[-1].item()"
# Default softmax scale
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(head_dim)
if not isinstance(sm_scale, torch.Tensor):
sm_scale_t = torch.tensor(sm_scale, dtype=torch.float32, device=q_c.device)
else:
sm_scale_t = sm_scale.to(dtype=torch.float32, device=q_c.device)
inv_ln2 = torch.tensor(1.0 / math.log(2.0), dtype=torch.float32, device=q_c.device)
# Allocate outputs
output_c = torch.empty((batch_size, num_qo_heads, head_dim), dtype=torch.bfloat16, device=q_c.device)
lse_c = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=q_c.device)
# Extract strides
stride_q_b, stride_q_h, stride_q_d = q_c.stride()
stride_k_p, stride_k_ps, stride_k_h, stride_k_d = k_c.stride()
stride_v_p, stride_v_ps, stride_v_h, stride_v_d = v_c.stride()
stride_out_b, stride_out_h, stride_out_d = output_c.stride()
stride_lse_b, stride_lse_h = lse_c.stride()
# Launch kernel
BLOCK_D = 128
STEP = 32
BLOCK_T = 128
GQA_RATIO = 4
grid = (batch_size * num_qo_heads,)
gqa_paged_decode_h32_kv8_d128_ps1_kernel[grid](
q_c, k_c, v_c,
kv_indptr_c, kv_indices_c,
output_c, lse_c,
sm_scale_t, inv_ln2,
batch_size,
stride_q_b, stride_q_h, stride_q_d,
stride_k_p, stride_k_ps, stride_k_h, stride_k_d,
stride_v_p, stride_v_ps, stride_v_h, stride_v_d,
stride_out_b, stride_out_h, stride_out_d,
stride_lse_b, stride_lse_h,
BLOCK_T=BLOCK_T, BLOCK_D=BLOCK_D, STEP=STEP, GQA_RATIO=GQA_RATIO,
num_warps=8, num_stages=3,
)
# Move outputs back to original device of q if needed
if output_c.device != device_out:
output = output_c.to(device_out)
lse = lse_c.to(device_out)
else:
output = output_c
lse = lse_c
return output, lsescrolls · 344 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON