claude-opus-4-1 / tritonc0a741
claude-opus-4-1_triton_c0a741 · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 251 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-c0a741?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:6ec51c2047845ee3afc5ed3133188b43417d850685e6dc65ad42b4898538ee7e
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 = 2
num_warps=2,stages = 1
num_stages=1Kernel source
main.py251 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def mla_paged_prefill_kernel_optimized(
q_nope_ptr, q_pe_ptr, ckv_cache_ptr, kpe_cache_ptr,
qo_indptr_ptr, kv_indptr_ptr, kv_indices_ptr,
output_ptr, lse_ptr,
sm_scale, total_q,
stride_qn_q, stride_qn_h, stride_qn_d,
stride_qp_q, stride_qp_h, stride_qp_d,
stride_ckv_p, stride_ckv_d,
stride_kpe_p, stride_kpe_d,
stride_o_q, stride_o_h, stride_o_d,
stride_lse_q, stride_lse_h,
batch_size,
BLOCK_D: tl.constexpr,
):
# Combined grid for all queries and heads
pid = tl.program_id(0)
num_heads = 16
# Compute query and head index
global_q_idx = pid // num_heads
head_idx = pid % num_heads
if global_q_idx >= total_q:
return
# Binary search for batch index
batch_idx = 0
left = 0
right = batch_size - 1
while left <= right:
mid = (left + right) // 2
q_start_mid = tl.load(qo_indptr_ptr + mid)
q_end_mid = tl.load(qo_indptr_ptr + mid + 1)
if global_q_idx < q_start_mid:
right = mid - 1
elif global_q_idx >= q_end_mid:
left = mid + 1
else:
batch_idx = mid
break
# Load batch 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)
q_len = q_end - q_start
kv_len = kv_end - kv_start
if kv_len <= 0:
# Store zeros for empty sequences
out_base = output_ptr + global_q_idx * stride_o_q + head_idx * stride_o_h
d_range = tl.arange(0, BLOCK_D)
zeros = tl.zeros([BLOCK_D], dtype=tl.bfloat16)
for offset in range(0, 512, BLOCK_D):
tl.store(out_base + (d_range + offset) * stride_o_d, zeros, mask=(d_range + offset) < 512)
tl.store(lse_ptr + global_q_idx * stride_lse_q + head_idx * stride_lse_h, float('-inf'))
return
q_idx = global_q_idx - q_start
# Causal mask computation
prefix_len = kv_len - q_len
query_abs_pos = prefix_len + q_idx
# Load query vectors
q_nope_base = q_nope_ptr + global_q_idx * stride_qn_q + head_idx * stride_qn_h
q_pe_base = q_pe_ptr + global_q_idx * stride_qp_q + head_idx * stride_qp_h
# Load query pe (64 dims)
d_range = tl.arange(0, BLOCK_D)
q_pe = tl.load(q_pe_base + d_range * stride_qp_d, mask=d_range < 64, other=0.0).to(tl.float32)
# Load query nope in blocks
q_blocks = []
for offset in range(0, 512, BLOCK_D):
q_block = tl.load(q_nope_base + (d_range + offset) * stride_qn_d, mask=(d_range + offset) < 512).to(tl.float32)
q_blocks.append(q_block)
# Initialize accumulators
max_logit = float('-inf')
sum_exp = 0.0
acc_blocks = []
for _ in range(8):
acc_blocks.append(tl.zeros([BLOCK_D], dtype=tl.float32))
# Process KV tokens one by one to reduce memory usage
for kv_idx in range(kv_len):
# Apply causal mask
if kv_idx > query_abs_pos:
break
# Get page index for this position
page_idx = tl.load(kv_indices_ptr + kv_start + kv_idx)
# Load key vectors
kc_base = ckv_cache_ptr + page_idx * stride_ckv_p
kp_base = kpe_cache_ptr + page_idx * stride_kpe_p
# Load kpe
kp = tl.load(kp_base + d_range * stride_kpe_d, mask=d_range < 64, other=0.0).to(tl.float32)
# Compute score
score = tl.sum(q_pe * kp)
# Load kc in blocks and compute dot product
kc_blocks = []
for i, offset in enumerate(range(0, 512, BLOCK_D)):
kc_block = tl.load(kc_base + (d_range + offset) * stride_ckv_d, mask=(d_range + offset) < 512).to(tl.float32)
kc_blocks.append(kc_block)
score += tl.sum(q_blocks[i] * kc_block)
score *= sm_scale
# Online softmax
if score > max_logit:
if max_logit > float('-inf'):
scale = tl.exp(max_logit - score)
sum_exp *= scale
for i in range(8):
acc_blocks[i] *= scale
max_logit = score
exp_score = 1.0
else:
exp_score = tl.exp(score - max_logit)
sum_exp += exp_score
# Accumulate
for i in range(8):
acc_blocks[i] += exp_score * kc_blocks[i]
# Store output
out_base = output_ptr + global_q_idx * stride_o_q + head_idx * stride_o_h
if sum_exp > 0:
inv_sum = 1.0 / sum_exp
for i, offset in enumerate(range(0, 512, BLOCK_D)):
result = (acc_blocks[i] * inv_sum).to(tl.bfloat16)
tl.store(out_base + (d_range + offset) * stride_o_d, result, mask=(d_range + offset) < 512)
# Store LSE in log base 2
log2_e = 1.44269504089
lse_val = (max_logit + tl.log(sum_exp)) * log2_e
else:
zeros = tl.zeros([BLOCK_D], dtype=tl.bfloat16)
for offset in range(0, 512, BLOCK_D):
tl.store(out_base + (d_range + offset) * stride_o_d, zeros, mask=(d_range + offset) < 512)
lse_val = float('-inf')
tl.store(lse_ptr + global_q_idx * stride_lse_q + head_idx * stride_lse_h, lse_val)
def run(q_nope, q_pe, ckv_cache, kpe_cache, qo_indptr, kv_indptr, kv_indices, sm_scale):
# Handle device placement
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available. This kernel requires GPU.")
# Store original devices
original_devices = {
'q_nope': q_nope.device,
'q_pe': q_pe.device,
'output': q_nope.device,
'lse': q_nope.device
}
# Move all tensors to GPU if needed
if not q_nope.is_cuda:
q_nope = q_nope.cuda()
if not q_pe.is_cuda:
q_pe = q_pe.cuda()
if not ckv_cache.is_cuda:
ckv_cache = ckv_cache.cuda()
if not kpe_cache.is_cuda:
kpe_cache = kpe_cache.cuda()
if not qo_indptr.is_cuda:
qo_indptr = qo_indptr.cuda()
if not kv_indptr.is_cuda:
kv_indptr = kv_indptr.cuda()
if not kv_indices.is_cuda:
kv_indices = kv_indices.cuda()
device = q_nope.device
# Get dimensions
total_q, num_qo_heads, head_dim_ckv = q_nope.shape
head_dim_kpe = q_pe.shape[-1]
page_size = ckv_cache.shape[1]
num_pages = ckv_cache.shape[0]
len_indptr = qo_indptr.shape[0]
batch_size = len_indptr - 1
num_kv_indices = kv_indices.shape[0]
# Verify constants
assert num_qo_heads == 16
assert head_dim_ckv == 512
assert head_dim_kpe == 64
assert page_size == 1
# Initialize outputs
output = torch.zeros((total_q, num_qo_heads, head_dim_ckv), dtype=torch.bfloat16, device=device)
lse = torch.full((total_q, num_qo_heads), float('-inf'), dtype=torch.float32, device=device)
# Reshape caches for page_size=1
ckv_cache = ckv_cache.squeeze(1) # [num_pages, head_dim_ckv]
kpe_cache = kpe_cache.squeeze(1) # [num_pages, head_dim_kpe]
# Check if there's any work to do
if batch_size == 0 or total_q == 0:
if not original_devices['output'].type == 'cuda':
output = output.cpu()
if not original_devices['lse'].type == 'cuda':
lse = lse.cpu()
return output, lse
# Use optimized kernel with combined grid
BLOCK_D = 64
grid = (total_q * num_qo_heads,)
mla_paged_prefill_kernel_optimized[grid](
q_nope, q_pe, ckv_cache, kpe_cache,
qo_indptr, kv_indptr, kv_indices,
output, lse,
sm_scale, total_q,
q_nope.stride(0), q_nope.stride(1), q_nope.stride(2),
q_pe.stride(0), q_pe.stride(1), q_pe.stride(2),
ckv_cache.stride(0), ckv_cache.stride(1),
kpe_cache.stride(0), kpe_cache.stride(1),
output.stride(0), output.stride(1), output.stride(2),
lse.stride(0), lse.stride(1),
batch_size,
BLOCK_D=BLOCK_D,
num_warps=2,
num_stages=1
)
# Move results back to original devices if needed
if not original_devices['output'].type == 'cuda':
output = output.cpu()
if not original_devices['lse'].type == 'cuda':
lse = lse.cpu()
return output, lsescrolls · 251 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON