claude-opus-4-1_triton_b32529
claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 182 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-b32529?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
2 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:9d7572352b293513002c2fdeaecd07c16d62ff8cd03bc8cd3b8b2db2813ae9e1
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.
Kernel source
main.py182 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def gqa_paged_prefill_kernel(
q_ptr, k_cache_ptr, v_cache_ptr,
qo_indptr_ptr, kv_indptr_ptr, kv_indices_ptr,
output_ptr, lse_ptr,
sm_scale,
q_start, q_end, kv_start, kv_end,
total_q, num_qo_heads, num_kv_heads, head_dim,
BLOCK_D: tl.constexpr,
):
# Get program IDs
pid_q = tl.program_id(0)
pid_h = tl.program_id(1)
# Compute indices
local_q_idx = pid_q
head_id = pid_h
num_q_tokens = q_end - q_start
num_kv_tokens = kv_end - kv_start
if local_q_idx >= num_q_tokens or head_id >= num_qo_heads:
return
global_q_idx = q_start + local_q_idx
# Find corresponding KV head for GQA
gqa_ratio = num_qo_heads // num_kv_heads
kv_head_id = head_id // gqa_ratio
# Delta for causal masking
delta = num_kv_tokens - num_q_tokens
max_kv_idx = tl.minimum(local_q_idx + 1 + delta, num_kv_tokens)
if max_kv_idx <= 0:
return
# Load query vector
d_offs = tl.arange(0, BLOCK_D)
q_offset = global_q_idx * num_qo_heads * head_dim + head_id * head_dim + d_offs
mask_d = d_offs < head_dim
q_vec = tl.load(q_ptr + q_offset, mask=mask_d, other=0.0).to(tl.float32)
# Initialize accumulators for online softmax
m_i = -float('inf')
l_i = 0.0
acc = tl.zeros([BLOCK_D], dtype=tl.float32)
# Process KV tokens one by one for better memory efficiency
for kv_idx in range(max_kv_idx):
# Load page ID
page_id = tl.load(kv_indices_ptr + kv_start + kv_idx)
# Load K vector
k_offset = page_id * num_kv_heads * head_dim + kv_head_id * head_dim + d_offs
k_vec = tl.load(k_cache_ptr + k_offset, mask=mask_d, other=0.0).to(tl.float32)
# Compute score
score = tl.sum(q_vec * k_vec, axis=0)
score = score * sm_scale
# Online softmax update
m_new = tl.maximum(m_i, score)
exp_score = tl.exp(score - m_new)
exp_m_diff = tl.exp(m_i - m_new)
# Update running sum
l_new = exp_m_diff * l_i + exp_score
# Rescale accumulator
acc = acc * exp_m_diff
# Load V vector and accumulate
v_offset = page_id * num_kv_heads * head_dim + kv_head_id * head_dim + d_offs
v_vec = tl.load(v_cache_ptr + v_offset, mask=mask_d, other=0.0).to(tl.float32)
acc = acc + v_vec * exp_score
# Update max and sum
m_i = m_new
l_i = l_new
# Normalize and store output
if l_i > 0:
output_vec = (acc / l_i).to(tl.bfloat16)
out_offset = global_q_idx * num_qo_heads * head_dim + head_id * head_dim + d_offs
tl.store(output_ptr + out_offset, output_vec, mask=mask_d)
# Store LSE (convert to base 2)
log2_e = 1.4426950408889634
lse_val = (m_i + tl.log(l_i)) * log2_e
lse_offset = global_q_idx * num_qo_heads + head_id
tl.store(lse_ptr + lse_offset, lse_val)
def run(q, k_cache, v_cache, qo_indptr, kv_indptr, kv_indices, sm_scale):
# Store original devices
original_device = q.device
# Move to GPU if needed
if not q.is_cuda:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available for GPU tensors")
device = torch.device('cuda')
q = q.cuda()
k_cache = k_cache.cuda()
v_cache = v_cache.cuda()
qo_indptr = qo_indptr.cuda()
kv_indptr = kv_indptr.cuda()
kv_indices = kv_indices.cuda()
else:
device = q.device
# Ensure all tensors are on same device
k_cache = k_cache.to(device)
v_cache = v_cache.to(device)
qo_indptr = qo_indptr.to(device)
kv_indptr = kv_indptr.to(device)
kv_indices = kv_indices.to(device)
# Get dimensions
total_q, num_qo_heads, head_dim = q.shape
num_pages, page_size, num_kv_heads, _ = k_cache.shape
len_indptr = qo_indptr.shape[0]
# Constants
assert num_qo_heads == 32
assert num_kv_heads == 8
assert head_dim == 128
assert page_size == 1
# Allocate outputs on device
output = torch.zeros((total_q, num_qo_heads, head_dim), dtype=torch.bfloat16, device=device)
lse = torch.full((total_q, num_qo_heads), -float('inf'), dtype=torch.float32, device=device)
# Flatten cache tensors since page_size=1
k_cache_flat = k_cache.squeeze(1) # [num_pages, num_kv_heads, head_dim]
v_cache_flat = v_cache.squeeze(1) # [num_pages, num_kv_heads, head_dim]
# Process each batch
num_batches = len_indptr - 1
# Choose block sizes
BLOCK_D = 128 # Since head_dim is 128
for batch_id in range(num_batches):
q_start = qo_indptr[batch_id].item()
q_end = qo_indptr[batch_id + 1].item()
kv_start = kv_indptr[batch_id].item()
kv_end = kv_indptr[batch_id + 1].item()
if q_start >= q_end or kv_start >= kv_end:
continue
num_q_tokens = q_end - q_start
# Use 2D grid for better parallelization
grid = (num_q_tokens, num_qo_heads)
gqa_paged_prefill_kernel[grid](
q, k_cache_flat, v_cache_flat,
qo_indptr, kv_indptr, kv_indices,
output, lse,
sm_scale,
q_start, q_end, kv_start, kv_end,
total_q, num_qo_heads, num_kv_heads, head_dim,
BLOCK_D=BLOCK_D,
num_warps=4,
num_stages=2,
)
# Move outputs back to original device if needed
if not original_device.type == 'cuda':
output = output.cpu()
lse = lse.cpu()
elif original_device != device:
output = output.to(original_device)
lse = lse.to(original_device)
return output, lsescrolls · 182 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON