claude-opus-4-1_triton_a98005
claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 207 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-a98005?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
No published measurement for this revision.
No evidence · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:b850cb66b93e401873ac5b58f5be5ed2353fa59fdc8dc239d45266c45a3372ad
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.
num-warps = 4
num_warps=4,stages = 1
num_stages=1,Kernel source
main.py207 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def mla_paged_decode_kernel(
q_nope_ptr, q_pe_ptr, ckv_cache_ptr, kpe_cache_ptr,
kv_indptr_ptr, kv_indices_ptr,
output_ptr, lse_ptr,
sm_scale,
batch_size,
stride_qn_b, stride_qn_h,
stride_qp_b, stride_qp_h,
stride_o_b, stride_o_h,
stride_lse_b,
HEAD_DIM_CKV: tl.constexpr,
HEAD_DIM_KPE: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
batch_idx = tl.program_id(0)
head_idx = tl.program_id(1)
if batch_idx >= batch_size:
return
# Get KV range for this batch
kv_start = tl.load(kv_indptr_ptr + batch_idx)
kv_end = tl.load(kv_indptr_ptr + batch_idx + 1)
kv_len = kv_end - kv_start
if kv_len <= 0:
# Write zeros for empty sequences
for d_offset in range(0, HEAD_DIM_CKV, BLOCK_SIZE):
d_range = tl.arange(0, BLOCK_SIZE) + d_offset
mask = d_range < HEAD_DIM_CKV
output_offset = batch_idx * stride_o_b + head_idx * stride_o_h + d_range
tl.store(output_ptr + output_offset, tl.zeros((BLOCK_SIZE,), dtype=tl.bfloat16), mask=mask)
lse_offset = batch_idx * stride_lse_b + head_idx
tl.store(lse_ptr + lse_offset, -float('inf'))
return
# Load query vectors
q_base = batch_idx * stride_qn_b + head_idx * stride_qn_h
qp_base = batch_idx * stride_qp_b + head_idx * stride_qp_h
# Load q_nope in chunks
qn_chunks = []
num_chunks = HEAD_DIM_CKV // BLOCK_SIZE
for i in range(num_chunks):
offset = i * BLOCK_SIZE
d_range = tl.arange(0, BLOCK_SIZE) + offset
qn_chunk = tl.load(q_nope_ptr + q_base + d_range).to(tl.float32)
qn_chunks.append(qn_chunk)
# Load q_pe
qp_range = tl.arange(0, HEAD_DIM_KPE)
qp = tl.load(q_pe_ptr + qp_base + qp_range).to(tl.float32)
# Initialize accumulators
max_logit = -float('inf')
sum_exp = 0.0
acc_chunks = []
for i in range(num_chunks):
acc_chunks.append(tl.zeros((BLOCK_SIZE,), dtype=tl.float32))
# Process KV tokens in blocks for better memory efficiency
KV_BLOCK = 32
for kv_block_start in range(0, kv_len, KV_BLOCK):
kv_block_end = tl.minimum(kv_block_start + KV_BLOCK, kv_len)
for kv_offset in range(kv_block_end - kv_block_start):
kv_idx = kv_block_start + kv_offset
if kv_idx >= kv_len:
break
page_idx = tl.load(kv_indices_ptr + kv_start + kv_idx)
kc_base = page_idx * HEAD_DIM_CKV
kp_base = page_idx * HEAD_DIM_KPE
# Compute dot product for q_nope and ckv
dot_nope = 0.0
kc_chunks = []
for i in range(num_chunks):
offset = i * BLOCK_SIZE
d_range = tl.arange(0, BLOCK_SIZE) + offset
kc_chunk = tl.load(ckv_cache_ptr + kc_base + d_range).to(tl.float32)
kc_chunks.append(kc_chunk)
dot_nope += tl.sum(qn_chunks[i] * kc_chunk)
# Compute dot product for q_pe and kpe
kp = tl.load(kpe_cache_ptr + kp_base + qp_range).to(tl.float32)
dot_pe = tl.sum(qp * kp)
# Compute scaled logit
logit = (dot_nope + dot_pe) * sm_scale
# Online softmax update
new_max = tl.maximum(max_logit, logit)
# Rescale previous accumulator
if kv_idx > 0 and max_logit > -float('inf'):
scale = tl.exp(max_logit - new_max)
sum_exp *= scale
for i in range(num_chunks):
acc_chunks[i] *= scale
max_logit = new_max
exp_val = tl.exp(logit - max_logit)
sum_exp += exp_val
# Accumulate weighted kc
for i in range(num_chunks):
acc_chunks[i] += exp_val * kc_chunks[i]
# Write output
o_base = batch_idx * stride_o_b + head_idx * stride_o_h
if sum_exp > 0:
inv_sum = 1.0 / sum_exp
for i in range(num_chunks):
offset = i * BLOCK_SIZE
d_range = tl.arange(0, BLOCK_SIZE) + offset
tl.store(output_ptr + o_base + d_range,
(acc_chunks[i] * inv_sum).to(tl.bfloat16))
else:
for i in range(num_chunks):
offset = i * BLOCK_SIZE
d_range = tl.arange(0, BLOCK_SIZE) + offset
tl.store(output_ptr + o_base + d_range,
tl.zeros((BLOCK_SIZE,), dtype=tl.bfloat16))
# Compute and store LSE (2-based)
lse_val = -float('inf')
if sum_exp > 0 and max_logit > -float('inf'):
log2_e = 1.0 / math.log(2.0)
lse_val = (max_logit + tl.log(sum_exp)) * log2_e
lse_offset = batch_idx * stride_lse_b + head_idx
tl.store(lse_ptr + lse_offset, lse_val)
def run(q_nope, q_pe, ckv_cache, kpe_cache, kv_indptr, kv_indices, sm_scale):
# Device management
device = q_nope.device
original_device = device
# Move to GPU if needed
if device.type == 'cpu':
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available for GPU computation")
device = torch.device('cuda')
q_nope = q_nope.cuda()
q_pe = q_pe.cuda()
ckv_cache = ckv_cache.cuda()
kpe_cache = kpe_cache.cuda()
kv_indptr = kv_indptr.cuda()
kv_indices = kv_indices.cuda()
batch_size, num_qo_heads, head_dim_ckv = q_nope.shape
head_dim_kpe = q_pe.shape[-1]
# Squeeze out page_size dimension (=1)
ckv_cache_flat = ckv_cache.squeeze(1).contiguous()
kpe_cache_flat = kpe_cache.squeeze(1).contiguous()
# Make inputs contiguous
q_nope = q_nope.contiguous()
q_pe = q_pe.contiguous()
kv_indptr = kv_indptr.contiguous()
kv_indices = kv_indices.contiguous()
# Allocate outputs
output = torch.zeros((batch_size, num_qo_heads, head_dim_ckv),
dtype=torch.bfloat16, device=device)
lse = torch.full((batch_size, num_qo_heads), -float('inf'),
dtype=torch.float32, device=device)
# Launch kernel with optimized configuration
grid = (batch_size, num_qo_heads)
# Use smaller block size to reduce memory usage
BLOCK_SIZE = 64
mla_paged_decode_kernel[grid](
q_nope, q_pe, ckv_cache_flat, kpe_cache_flat,
kv_indptr, kv_indices,
output, lse,
sm_scale,
batch_size,
q_nope.stride(0), q_nope.stride(1),
q_pe.stride(0), q_pe.stride(1),
output.stride(0), output.stride(1),
lse.stride(0),
HEAD_DIM_CKV=head_dim_ckv,
HEAD_DIM_KPE=head_dim_kpe,
BLOCK_SIZE=BLOCK_SIZE,
num_warps=4,
num_stages=1,
)
# Move back to original device if needed
if original_device.type == 'cpu':
output = output.cpu()
lse = lse.cpu()
return output, lsescrolls · 207 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON