Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / triton07ad16

claude-opus-4-1_triton_07ad16 · claude-opus-4-1-20250805 · triton · Apache-2.0

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
NVIDIA B200
100.3µs
#1 of 6
2025-10-21

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

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

Best evidence level for this revision: reported

JSON