claude-opus-4-1 / triton4080e2
claude-opus-4-1_triton_4080e2 · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 323 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-4080e2?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=18 · num_kv_indices=2
NVIDIA B200
14.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=11 · num_kv_indices=10
NVIDIA B200
15.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=63 · num_kv_indices=46
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
15.5µs
#3 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=55 · num_kv_indices=38
NVIDIA B200
15.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
15.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
16.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
20.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
26.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1396 · num_kv_indices=1369
NVIDIA B200
29.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
29.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
30.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1732 · num_kv_indices=1705
NVIDIA B200
30.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
30.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
34.0µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
34.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2212 · num_kv_indices=2185
NVIDIA B200
34.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
35.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
35.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2708 · num_kv_indices=2681
NVIDIA B200
35.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
36.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
39.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=4333 · num_kv_indices=4298
NVIDIA B200
45.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
59.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
112.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
151.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
243.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
248.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
248.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63437 · num_kv_indices=54166
NVIDIA B200
249.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
250.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64205 · num_kv_indices=54934
NVIDIA B200
251.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
252.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
253.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65037 · num_kv_indices=55766
NVIDIA B200
254.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
255.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
255.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
256.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67021 · num_kv_indices=57750
NVIDIA B200
257.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67405 · num_kv_indices=58134
NVIDIA B200
260.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
263.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=68237 · num_kv_indices=58966
NVIDIA B200
267.4µs
#2 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:d399c1b37d3746d464e37c57904264b7d2a27e5a5c66b03dedf4f37bec7c8e73
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
online-softmax
m_i_new = tl.maximum(m_i, tl.max(logits, axis=0))tile-m = 128
BLOCK_M = 128 # Process more tokens per block for B200Kernel source
main.py323 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def gqa_paged_decode_kernel(
q_ptr, k_cache_ptr, v_cache_ptr,
kv_indptr_ptr, kv_indices_ptr,
output_ptr, lse_ptr,
sm_scale,
batch_size, num_pages,
BLOCK_M: tl.constexpr,
BLOCK_D: tl.constexpr,
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
GQA_RATIO: tl.constexpr,
):
# Grid indices
batch_idx = tl.program_id(0)
head_idx = tl.program_id(1)
if batch_idx >= batch_size:
return
# Get KV head index for this query head (GQA)
kv_head_idx = head_idx // GQA_RATIO
# Get sequence bounds from indptr
seq_start = tl.load(kv_indptr_ptr + batch_idx)
seq_end = tl.load(kv_indptr_ptr + batch_idx + 1)
seq_len = seq_end - seq_start
# Calculate output offset once
output_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
d_idx = tl.arange(0, HEAD_DIM)
if seq_len <= 0:
# No KV cache for this batch element - write zeros
zeros = tl.zeros((HEAD_DIM,), dtype=tl.bfloat16)
tl.store(output_ptr + output_offset + d_idx, zeros)
lse_offset = batch_idx * NUM_QO_HEADS + head_idx
tl.store(lse_ptr + lse_offset, float('-inf'))
return
# Load query vector for this head
q_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
q = tl.load(q_ptr + q_offset + d_idx).to(tl.float32)
# Initialize accumulators
m_i = float('-inf') # Max logit
l_i = 0.0 # Sum of exponentials
acc = tl.zeros((HEAD_DIM,), dtype=tl.float32)
# Process KV cache tokens in blocks
for token_start in range(0, seq_len, BLOCK_M):
token_range = tl.arange(0, BLOCK_M)
token_idx = token_start + token_range
token_mask = token_idx < seq_len
# Get page indices for this block of tokens
global_token_idx = seq_start + token_idx
page_idx = tl.load(kv_indices_ptr + global_token_idx, mask=token_mask, other=0)
# Initialize logits for this block
logits = tl.zeros((BLOCK_M,), dtype=tl.float32)
# Compute dot products efficiently
for d_start in range(0, HEAD_DIM, BLOCK_D):
d_range = tl.arange(0, BLOCK_D) + d_start
d_mask = d_range < HEAD_DIM
# Load K values for all tokens in block
k_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
k_offsets = k_base + d_range[None, :]
k_vals = tl.load(k_cache_ptr + k_offsets,
mask=token_mask[:, None] & d_mask[None, :],
other=0.0).to(tl.float32)
# Get query chunk using masking instead of slicing
q_chunk = tl.load(q_ptr + q_offset + d_range, mask=d_mask, other=0.0).to(tl.float32)
# Accumulate partial dot products
partial_dots = tl.sum(k_vals * q_chunk[None, :], axis=1)
logits += partial_dots
# Scale logits
logits = logits * sm_scale
logits = tl.where(token_mask, logits, float('-inf'))
# Online softmax: update running max and sum
m_i_new = tl.maximum(m_i, tl.max(logits, axis=0))
# Compute exponentials with numerical stability
exp_logits = tl.exp(logits - m_i_new)
exp_logits = tl.where(token_mask, exp_logits, 0.0)
# Update running sum with correction factor
correction = tl.exp(m_i - m_i_new)
l_i = l_i * correction + tl.sum(exp_logits, axis=0)
# Scale accumulator by correction factor
acc = acc * correction
# Update accumulator with V values
for d_start in range(0, HEAD_DIM, BLOCK_D):
d_range = tl.arange(0, BLOCK_D) + d_start
d_mask = d_range < HEAD_DIM
# Load V values for all tokens in block
v_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
v_offsets = v_base + d_range[None, :]
v_vals = tl.load(v_cache_ptr + v_offsets,
mask=token_mask[:, None] & d_mask[None, :],
other=0.0).to(tl.float32)
# Weighted sum with attention weights
weighted_v = tl.sum(v_vals * exp_logits[:, None], axis=0)
# Update the accumulator slice directly using masking
acc_update = tl.load(output_ptr + output_offset + d_range, mask=d_mask, other=0.0).to(tl.float32)
acc_slice = tl.load(q_ptr + q_offset + d_range, mask=d_mask, other=0.0).to(tl.float32) * 0.0 # Zero placeholder
# Store weighted_v to the corresponding positions in acc
for i in range(BLOCK_D):
if d_start + i < HEAD_DIM:
idx = d_start + i
mask_i = d_idx == idx
acc = tl.where(mask_i, acc + weighted_v[i], acc)
m_i = m_i_new
# Normalize and store output
output_val = (acc / l_i).to(tl.bfloat16)
tl.store(output_ptr + output_offset + d_idx, output_val)
# Store LSE (log-sum-exp in base 2)
lse_offset = batch_idx * NUM_QO_HEADS + head_idx
log2 = 0.6931471805599453 # math.log(2)
lse_val = (m_i + tl.log(l_i)) / log2
tl.store(lse_ptr + lse_offset, lse_val)
@triton.jit
def gqa_paged_decode_kernel_optimized(
q_ptr, k_cache_ptr, v_cache_ptr,
kv_indptr_ptr, kv_indices_ptr,
output_ptr, lse_ptr,
sm_scale,
batch_size, num_pages,
BLOCK_M: tl.constexpr,
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
GQA_RATIO: tl.constexpr,
):
# Grid indices
batch_idx = tl.program_id(0)
head_idx = tl.program_id(1)
if batch_idx >= batch_size:
return
# Get KV head index for this query head (GQA)
kv_head_idx = head_idx // GQA_RATIO
# Get sequence bounds from indptr
seq_start = tl.load(kv_indptr_ptr + batch_idx)
seq_end = tl.load(kv_indptr_ptr + batch_idx + 1)
seq_len = seq_end - seq_start
# Calculate output offset
output_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
if seq_len <= 0:
# No KV cache for this batch element - write zeros
d_idx = tl.arange(0, HEAD_DIM)
zeros = tl.zeros((HEAD_DIM,), dtype=tl.bfloat16)
tl.store(output_ptr + output_offset + d_idx, zeros)
lse_offset = batch_idx * NUM_QO_HEADS + head_idx
tl.store(lse_ptr + lse_offset, float('-inf'))
return
# Load entire query vector for this head
q_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
q_idx = tl.arange(0, HEAD_DIM)
q = tl.load(q_ptr + q_offset + q_idx).to(tl.float32)
# Initialize accumulators
m_i = float('-inf') # Max logit
l_i = 0.0 # Sum of exponentials
acc = tl.zeros((HEAD_DIM,), dtype=tl.float32)
# Process KV cache tokens in blocks
for token_start in range(0, seq_len, BLOCK_M):
token_range = tl.arange(0, BLOCK_M)
token_idx = token_start + token_range
token_mask = token_idx < seq_len
# Get page indices for this block of tokens
global_token_idx = seq_start + token_idx
page_idx = tl.load(kv_indices_ptr + global_token_idx, mask=token_mask, other=0)
# Compute dot products for all tokens in block at once
# Load K values and compute dot products
k_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
d_idx_expanded = tl.arange(0, HEAD_DIM)[None, :]
k_offsets = k_base + d_idx_expanded
k_vals = tl.load(k_cache_ptr + k_offsets,
mask=token_mask[:, None],
other=0.0).to(tl.float32)
# Compute dot products
logits = tl.sum(k_vals * q[None, :], axis=1)
# Scale logits
logits = logits * sm_scale
logits = tl.where(token_mask, logits, float('-inf'))
# Online softmax: update running max and sum
m_i_new = tl.maximum(m_i, tl.max(logits, axis=0))
# Compute exponentials with numerical stability
exp_logits = tl.exp(logits - m_i_new)
exp_logits = tl.where(token_mask, exp_logits, 0.0)
# Update running sum with correction factor
correction = tl.exp(m_i - m_i_new)
l_i = l_i * correction + tl.sum(exp_logits, axis=0)
# Scale accumulator by correction factor
acc = acc * correction
# Load V values and accumulate
v_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
v_offsets = v_base + d_idx_expanded
v_vals = tl.load(v_cache_ptr + v_offsets,
mask=token_mask[:, None],
other=0.0).to(tl.float32)
# Weighted sum with attention weights
weighted_v = tl.sum(v_vals * exp_logits[:, None], axis=0)
acc = acc + weighted_v
m_i = m_i_new
# Normalize and store output
output_val = (acc / l_i).to(tl.bfloat16)
tl.store(output_ptr + output_offset + q_idx, output_val)
# Store LSE (log-sum-exp in base 2)
lse_offset = batch_idx * NUM_QO_HEADS + head_idx
log2 = 0.6931471805599453 # math.log(2)
lse_val = (m_i + tl.log(l_i)) / log2
tl.store(lse_ptr + lse_offset, lse_val)
def run(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale):
# Handle device management
device = None
if q.is_cuda:
device = q.device
else:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU tensors are required")
device = torch.device('cuda')
q = q.cuda()
# Move all tensors to same device if needed
if not k_cache.is_cuda:
k_cache = k_cache.to(device)
if not v_cache.is_cuda:
v_cache = v_cache.to(device)
if not kv_indptr.is_cuda:
kv_indptr = kv_indptr.to(device)
if not kv_indices.is_cuda:
kv_indices = kv_indices.to(device)
# Get dimensions
batch_size, num_qo_heads, head_dim = q.shape
num_pages, page_size, num_kv_heads, _ = k_cache.shape
# Verify constants
assert num_qo_heads == 32, f"num_qo_heads must be 32, got {num_qo_heads}"
assert num_kv_heads == 8, f"num_kv_heads must be 8, got {num_kv_heads}"
assert head_dim == 128, f"head_dim must be 128, got {head_dim}"
assert page_size == 1, f"page_size must be 1, got {page_size}"
# GQA ratio
gqa_ratio = num_qo_heads // num_kv_heads
# Allocate outputs
output = torch.zeros((batch_size, num_qo_heads, head_dim), dtype=torch.bfloat16, device=device)
lse = torch.full((batch_size, num_qo_heads), -float('inf'), dtype=torch.float32, device=device)
# Flatten k_cache and v_cache for page_size=1
# Shape: [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, head_dim]
k_cache_flat = k_cache.squeeze(1)
v_cache_flat = v_cache.squeeze(1)
# Configure kernel - optimized for B200
BLOCK_M = 128 # Process more tokens per block for B200
# Launch kernel
grid = (batch_size, num_qo_heads)
gqa_paged_decode_kernel_optimized[grid](
q, k_cache_flat, v_cache_flat,
kv_indptr, kv_indices,
output, lse,
sm_scale,
batch_size, num_pages,
BLOCK_M=BLOCK_M,
NUM_QO_HEADS=num_qo_heads,
NUM_KV_HEADS=num_kv_heads,
HEAD_DIM=head_dim,
GQA_RATIO=gqa_ratio,
)
return output, lsescrolls · 323 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON