claude-opus-4-1 / triton07ad16
claude-opus-4-1_triton_07ad16 · claude-opus-4-1-20250805 · triton · Apache-2.0
Kernel source · 208 lines ↓holds 1 record
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 208 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-07ad16?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
1 measurement 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:28b5f8cac120829849007bc923430c0ddbac70ef0fc350a13012ea6af6f75071
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py208 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def gqa_paged_prefill_causal_kernel(
q_ptr, k_cache_ptr, v_cache_ptr,
qo_indptr_ptr, kv_indptr_ptr, kv_indices_ptr,
output_ptr, lse_ptr,
sm_scale,
total_q, num_pages, len_indptr, num_kv_indices,
BLOCK_KV: tl.constexpr,
HEAD_DIM: tl.constexpr,
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
GQA_RATIO: tl.constexpr,
):
# Grid: (batch_idx, q_head_idx, q_token_idx)
batch_idx = tl.program_id(0)
q_head_idx = tl.program_id(1)
q_token_idx = tl.program_id(2)
# Early exit for invalid batch
if batch_idx >= len_indptr - 1:
return
# Load sequence boundaries
q_start = tl.load(qo_indptr_ptr + batch_idx)
q_end = tl.load(qo_indptr_ptr + batch_idx + 1)
kv_start = tl.load(kv_indptr_ptr + batch_idx)
kv_end = tl.load(kv_indptr_ptr + batch_idx + 1)
num_q_tokens = q_end - q_start
num_kv_tokens = kv_end - kv_start
# Check if this q_token_idx is valid for this batch
if q_token_idx >= num_q_tokens:
return
if num_q_tokens <= 0 or num_kv_tokens <= 0:
return
global_q_idx = q_start + q_token_idx
# Causal mask limit
delta = num_kv_tokens - num_q_tokens
max_kv_idx = tl.minimum(q_token_idx + 1 + delta, num_kv_tokens)
# Skip if no valid KV tokens
if max_kv_idx <= 0:
return
# Determine KV head for this query head (GQA)
kv_head_idx = q_head_idx // GQA_RATIO
# Load query vector
q_offset = global_q_idx * NUM_QO_HEADS * HEAD_DIM + q_head_idx * HEAD_DIM
q_range = tl.arange(0, HEAD_DIM)
q = tl.load(q_ptr + q_offset + q_range).to(tl.float32)
# Initialize accumulators
numerator = tl.zeros([HEAD_DIM], dtype=tl.float32)
max_logit = -float('inf')
denominator = 0.0
# Process KV tokens in blocks
for kv_block_start in range(0, max_kv_idx, BLOCK_KV):
kv_block_end = tl.minimum(kv_block_start + BLOCK_KV, max_kv_idx)
kv_block_range = tl.arange(0, BLOCK_KV)
kv_mask = (kv_block_start + kv_block_range) < kv_block_end
# Load page indices for this block
kv_indices_offset = kv_start + kv_block_start
page_ids = tl.load(
kv_indices_ptr + kv_indices_offset + kv_block_range,
mask=kv_mask,
other=0
)
# Process each KV token in the block
logits = tl.zeros([BLOCK_KV], dtype=tl.float32)
# Compute logits for the block
for i in range(BLOCK_KV):
if kv_block_start + i < kv_block_end:
page_id = tl.load(kv_indices_ptr + kv_indices_offset + i)
# Load K vector from cache
k_offset = page_id * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
k = tl.load(k_cache_ptr + k_offset + q_range).to(tl.float32)
# Compute dot product
logit = tl.sum(q * k, axis=0) * sm_scale
logits = tl.where(tl.arange(0, BLOCK_KV) == i, logit, logits)
# Update max for numerical stability
block_max = tl.max(tl.where(kv_mask, logits, -float('inf')))
max_logit = tl.maximum(max_logit, block_max)
# Second pass: compute softmax and weighted sum with stable computation
for kv_block_start in range(0, max_kv_idx, BLOCK_KV):
kv_block_end = tl.minimum(kv_block_start + BLOCK_KV, max_kv_idx)
# Process each KV token in the block
for i in range(BLOCK_KV):
if kv_block_start + i < kv_block_end:
kv_indices_offset = kv_start + kv_block_start + i
page_id = tl.load(kv_indices_ptr + kv_indices_offset)
# Load K vector
k_offset = page_id * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
k = tl.load(k_cache_ptr + k_offset + q_range).to(tl.float32)
# Compute attention score
logit = tl.sum(q * k, axis=0) * sm_scale
score = tl.exp(logit - max_logit)
# Load V vector
v_offset = page_id * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
v = tl.load(v_cache_ptr + v_offset + q_range).to(tl.float32)
# Accumulate
numerator = numerator + score * v
denominator = denominator + score
# Normalize and store output
output = numerator / denominator
output_offset = global_q_idx * NUM_QO_HEADS * HEAD_DIM + q_head_idx * HEAD_DIM
tl.store(output_ptr + output_offset + q_range, output.to(tl.bfloat16))
# Compute and store LSE (log-sum-exp in base 2)
log2 = 0.6931471805599453 # math.log(2.0)
lse_value = (max_logit + tl.log(denominator)) / log2
lse_offset = global_q_idx * NUM_QO_HEADS + q_head_idx
tl.store(lse_ptr + lse_offset, lse_value)
def run(q, k_cache, v_cache, qo_indptr, kv_indptr, kv_indices, sm_scale=None):
# Store original device
original_device = q.device
# Device management
if q.is_cuda:
device = q.device
elif torch.cuda.is_available():
device = torch.device('cuda')
q = q.cuda()
k_cache = k_cache.cuda() if not k_cache.is_cuda else k_cache
v_cache = v_cache.cuda() if not v_cache.is_cuda else v_cache
qo_indptr = qo_indptr.cuda() if not qo_indptr.is_cuda else qo_indptr
kv_indptr = kv_indptr.cuda() if not kv_indptr.is_cuda else kv_indptr
kv_indices = kv_indices.cuda() if not kv_indices.is_cuda else kv_indices
else:
raise RuntimeError("CUDA is not available but GPU tensors are required")
# Extract 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]
num_kv_indices = kv_indices.shape[0]
# Verify constants
assert num_qo_heads == 32, f"Expected num_qo_heads=32, got {num_qo_heads}"
assert num_kv_heads == 4, f"Expected num_kv_heads=4, got {num_kv_heads}"
assert head_dim == 128, f"Expected head_dim=128, got {head_dim}"
assert page_size == 1, f"Expected page_size=1, got {page_size}"
# Set default sm_scale
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(head_dim)
# Allocate outputs
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)
# Configure kernel
BLOCK_KV = 64
GQA_RATIO = num_qo_heads // num_kv_heads
# Compute max queries per batch
max_q_per_batch = 0
for i in range(len_indptr - 1):
q_start = qo_indptr[i].item()
q_end = qo_indptr[i + 1].item()
max_q_per_batch = max(max_q_per_batch, q_end - q_start)
# Launch kernel with 3D grid
grid = (len_indptr - 1, num_qo_heads, max_q_per_batch)
gqa_paged_prefill_causal_kernel[grid](
q, k_cache, v_cache,
qo_indptr, kv_indptr, kv_indices,
output, lse,
sm_scale,
total_q, num_pages, len_indptr, num_kv_indices,
BLOCK_KV=BLOCK_KV,
HEAD_DIM=head_dim,
NUM_QO_HEADS=num_qo_heads,
NUM_KV_HEADS=num_kv_heads,
GQA_RATIO=GQA_RATIO,
)
# Move outputs back to original device if needed
if output.device != original_device:
output = output.to(original_device)
lse = lse.to(original_device)
return output, lsescrolls · 208 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON