Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
2.36ms
#6 of 7
2025-10-21
NVIDIA B200
2.39ms
#6 of 7
2025-10-21

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.

num-warps = 4num_warps=4,
online-softmaxm_new = tl.maximum(m_i, score)
stages = 2num_stages=2,

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, lse
scrolls · 182 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON