Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / triton4080e2

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 323 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-4080e2?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32

Benchmark evidence

48 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
14.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=11 · num_kv_indices=10
NVIDIA B200
15.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=63 · num_kv_indices=46
NVIDIA B200
15.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
Show all 48 measurements ›
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=55 · num_kv_indices=38
NVIDIA B200
15.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
15.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
16.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
20.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
26.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1396 · num_kv_indices=1369
NVIDIA B200
29.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
29.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
30.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1732 · num_kv_indices=1705
NVIDIA B200
30.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
30.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
34.0µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
34.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2212 · num_kv_indices=2185
NVIDIA B200
34.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
35.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
35.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2708 · num_kv_indices=2681
NVIDIA B200
35.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
36.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
39.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=4333 · num_kv_indices=4298
NVIDIA B200
45.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
59.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
112.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
151.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
243.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
248.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
248.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63437 · num_kv_indices=54166
NVIDIA B200
249.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
250.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64205 · num_kv_indices=54934
NVIDIA B200
251.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
252.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
253.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65037 · num_kv_indices=55766
NVIDIA B200
254.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
255.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
255.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
256.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67021 · num_kv_indices=57750
NVIDIA B200
257.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67405 · num_kv_indices=58134
NVIDIA B200
260.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
263.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=68237 · num_kv_indices=58966
NVIDIA B200
267.4µs
#2 of 7
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:d399c1b37d3746d464e37c57904264b7d2a27e5a5c66b03dedf4f37bec7c8e73
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.

online-softmaxm_i_new = tl.maximum(m_i, tl.max(logits, axis=0))
tile-m = 128BLOCK_M = 128 # Process more tokens per block for B200

Kernel source

main.py323 lines
import torch
import triton
import triton.language as tl
import math

@triton.jit
def gqa_paged_decode_kernel(
    q_ptr, k_cache_ptr, v_cache_ptr,
    kv_indptr_ptr, kv_indices_ptr,
    output_ptr, lse_ptr,
    sm_scale,
    batch_size, num_pages,
    BLOCK_M: tl.constexpr,
    BLOCK_D: tl.constexpr,
    NUM_QO_HEADS: tl.constexpr,
    NUM_KV_HEADS: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    GQA_RATIO: tl.constexpr,
):
    # Grid indices
    batch_idx = tl.program_id(0)
    head_idx = tl.program_id(1)
    
    if batch_idx >= batch_size:
        return
    
    # Get KV head index for this query head (GQA)
    kv_head_idx = head_idx // GQA_RATIO
    
    # Get sequence bounds from indptr
    seq_start = tl.load(kv_indptr_ptr + batch_idx)
    seq_end = tl.load(kv_indptr_ptr + batch_idx + 1)
    seq_len = seq_end - seq_start
    
    # Calculate output offset once
    output_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
    d_idx = tl.arange(0, HEAD_DIM)
    
    if seq_len <= 0:
        # No KV cache for this batch element - write zeros
        zeros = tl.zeros((HEAD_DIM,), dtype=tl.bfloat16)
        tl.store(output_ptr + output_offset + d_idx, zeros)
        
        lse_offset = batch_idx * NUM_QO_HEADS + head_idx
        tl.store(lse_ptr + lse_offset, float('-inf'))
        return
    
    # Load query vector for this head
    q_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
    q = tl.load(q_ptr + q_offset + d_idx).to(tl.float32)
    
    # Initialize accumulators
    m_i = float('-inf')  # Max logit
    l_i = 0.0  # Sum of exponentials
    acc = tl.zeros((HEAD_DIM,), dtype=tl.float32)
    
    # Process KV cache tokens in blocks
    for token_start in range(0, seq_len, BLOCK_M):
        token_range = tl.arange(0, BLOCK_M)
        token_idx = token_start + token_range
        token_mask = token_idx < seq_len
        
        # Get page indices for this block of tokens
        global_token_idx = seq_start + token_idx
        page_idx = tl.load(kv_indices_ptr + global_token_idx, mask=token_mask, other=0)
        
        # Initialize logits for this block
        logits = tl.zeros((BLOCK_M,), dtype=tl.float32)
        
        # Compute dot products efficiently
        for d_start in range(0, HEAD_DIM, BLOCK_D):
            d_range = tl.arange(0, BLOCK_D) + d_start
            d_mask = d_range < HEAD_DIM
            
            # Load K values for all tokens in block
            k_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
            k_offsets = k_base + d_range[None, :]
            k_vals = tl.load(k_cache_ptr + k_offsets, 
                           mask=token_mask[:, None] & d_mask[None, :], 
                           other=0.0).to(tl.float32)
            
            # Get query chunk using masking instead of slicing
            q_chunk = tl.load(q_ptr + q_offset + d_range, mask=d_mask, other=0.0).to(tl.float32)
            
            # Accumulate partial dot products
            partial_dots = tl.sum(k_vals * q_chunk[None, :], axis=1)
            logits += partial_dots
        
        # Scale logits
        logits = logits * sm_scale
        logits = tl.where(token_mask, logits, float('-inf'))
        
        # Online softmax: update running max and sum
        m_i_new = tl.maximum(m_i, tl.max(logits, axis=0))
        
        # Compute exponentials with numerical stability
        exp_logits = tl.exp(logits - m_i_new)
        exp_logits = tl.where(token_mask, exp_logits, 0.0)
        
        # Update running sum with correction factor
        correction = tl.exp(m_i - m_i_new)
        l_i = l_i * correction + tl.sum(exp_logits, axis=0)
        
        # Scale accumulator by correction factor
        acc = acc * correction
        
        # Update accumulator with V values
        for d_start in range(0, HEAD_DIM, BLOCK_D):
            d_range = tl.arange(0, BLOCK_D) + d_start
            d_mask = d_range < HEAD_DIM
            
            # Load V values for all tokens in block
            v_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
            v_offsets = v_base + d_range[None, :]
            v_vals = tl.load(v_cache_ptr + v_offsets, 
                           mask=token_mask[:, None] & d_mask[None, :], 
                           other=0.0).to(tl.float32)
            
            # Weighted sum with attention weights
            weighted_v = tl.sum(v_vals * exp_logits[:, None], axis=0)
            
            # Update the accumulator slice directly using masking
            acc_update = tl.load(output_ptr + output_offset + d_range, mask=d_mask, other=0.0).to(tl.float32)
            acc_slice = tl.load(q_ptr + q_offset + d_range, mask=d_mask, other=0.0).to(tl.float32) * 0.0  # Zero placeholder
            
            # Store weighted_v to the corresponding positions in acc
            for i in range(BLOCK_D):
                if d_start + i < HEAD_DIM:
                    idx = d_start + i
                    mask_i = d_idx == idx
                    acc = tl.where(mask_i, acc + weighted_v[i], acc)
        
        m_i = m_i_new
    
    # Normalize and store output
    output_val = (acc / l_i).to(tl.bfloat16)
    tl.store(output_ptr + output_offset + d_idx, output_val)
    
    # Store LSE (log-sum-exp in base 2)
    lse_offset = batch_idx * NUM_QO_HEADS + head_idx
    log2 = 0.6931471805599453  # math.log(2)
    lse_val = (m_i + tl.log(l_i)) / log2
    tl.store(lse_ptr + lse_offset, lse_val)


@triton.jit
def gqa_paged_decode_kernel_optimized(
    q_ptr, k_cache_ptr, v_cache_ptr,
    kv_indptr_ptr, kv_indices_ptr,
    output_ptr, lse_ptr,
    sm_scale,
    batch_size, num_pages,
    BLOCK_M: tl.constexpr,
    NUM_QO_HEADS: tl.constexpr,
    NUM_KV_HEADS: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    GQA_RATIO: tl.constexpr,
):
    # Grid indices
    batch_idx = tl.program_id(0)
    head_idx = tl.program_id(1)
    
    if batch_idx >= batch_size:
        return
    
    # Get KV head index for this query head (GQA)
    kv_head_idx = head_idx // GQA_RATIO
    
    # Get sequence bounds from indptr
    seq_start = tl.load(kv_indptr_ptr + batch_idx)
    seq_end = tl.load(kv_indptr_ptr + batch_idx + 1)
    seq_len = seq_end - seq_start
    
    # Calculate output offset
    output_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
    
    if seq_len <= 0:
        # No KV cache for this batch element - write zeros
        d_idx = tl.arange(0, HEAD_DIM)
        zeros = tl.zeros((HEAD_DIM,), dtype=tl.bfloat16)
        tl.store(output_ptr + output_offset + d_idx, zeros)
        
        lse_offset = batch_idx * NUM_QO_HEADS + head_idx
        tl.store(lse_ptr + lse_offset, float('-inf'))
        return
    
    # Load entire query vector for this head
    q_offset = batch_idx * NUM_QO_HEADS * HEAD_DIM + head_idx * HEAD_DIM
    q_idx = tl.arange(0, HEAD_DIM)
    q = tl.load(q_ptr + q_offset + q_idx).to(tl.float32)
    
    # Initialize accumulators
    m_i = float('-inf')  # Max logit
    l_i = 0.0  # Sum of exponentials
    acc = tl.zeros((HEAD_DIM,), dtype=tl.float32)
    
    # Process KV cache tokens in blocks
    for token_start in range(0, seq_len, BLOCK_M):
        token_range = tl.arange(0, BLOCK_M)
        token_idx = token_start + token_range
        token_mask = token_idx < seq_len
        
        # Get page indices for this block of tokens
        global_token_idx = seq_start + token_idx
        page_idx = tl.load(kv_indices_ptr + global_token_idx, mask=token_mask, other=0)
        
        # Compute dot products for all tokens in block at once
        # Load K values and compute dot products
        k_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
        d_idx_expanded = tl.arange(0, HEAD_DIM)[None, :]
        k_offsets = k_base + d_idx_expanded
        k_vals = tl.load(k_cache_ptr + k_offsets, 
                        mask=token_mask[:, None], 
                        other=0.0).to(tl.float32)
        
        # Compute dot products
        logits = tl.sum(k_vals * q[None, :], axis=1)
        
        # Scale logits
        logits = logits * sm_scale
        logits = tl.where(token_mask, logits, float('-inf'))
        
        # Online softmax: update running max and sum
        m_i_new = tl.maximum(m_i, tl.max(logits, axis=0))
        
        # Compute exponentials with numerical stability
        exp_logits = tl.exp(logits - m_i_new)
        exp_logits = tl.where(token_mask, exp_logits, 0.0)
        
        # Update running sum with correction factor
        correction = tl.exp(m_i - m_i_new)
        l_i = l_i * correction + tl.sum(exp_logits, axis=0)
        
        # Scale accumulator by correction factor
        acc = acc * correction
        
        # Load V values and accumulate
        v_base = page_idx[:, None] * NUM_KV_HEADS * HEAD_DIM + kv_head_idx * HEAD_DIM
        v_offsets = v_base + d_idx_expanded
        v_vals = tl.load(v_cache_ptr + v_offsets, 
                        mask=token_mask[:, None], 
                        other=0.0).to(tl.float32)
        
        # Weighted sum with attention weights
        weighted_v = tl.sum(v_vals * exp_logits[:, None], axis=0)
        acc = acc + weighted_v
        
        m_i = m_i_new
    
    # Normalize and store output
    output_val = (acc / l_i).to(tl.bfloat16)
    tl.store(output_ptr + output_offset + q_idx, output_val)
    
    # Store LSE (log-sum-exp in base 2)
    lse_offset = batch_idx * NUM_QO_HEADS + head_idx
    log2 = 0.6931471805599453  # math.log(2)
    lse_val = (m_i + tl.log(l_i)) / log2
    tl.store(lse_ptr + lse_offset, lse_val)


def run(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale):
    # Handle device management
    device = None
    if q.is_cuda:
        device = q.device
    else:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU tensors are required")
        device = torch.device('cuda')
        q = q.cuda()
    
    # Move all tensors to same device if needed
    if not k_cache.is_cuda:
        k_cache = k_cache.to(device)
    if not v_cache.is_cuda:
        v_cache = v_cache.to(device)
    if not kv_indptr.is_cuda:
        kv_indptr = kv_indptr.to(device)
    if not kv_indices.is_cuda:
        kv_indices = kv_indices.to(device)
    
    # Get dimensions
    batch_size, num_qo_heads, head_dim = q.shape
    num_pages, page_size, num_kv_heads, _ = k_cache.shape
    
    # Verify constants
    assert num_qo_heads == 32, f"num_qo_heads must be 32, got {num_qo_heads}"
    assert num_kv_heads == 8, f"num_kv_heads must be 8, got {num_kv_heads}"
    assert head_dim == 128, f"head_dim must be 128, got {head_dim}"
    assert page_size == 1, f"page_size must be 1, got {page_size}"
    
    # GQA ratio
    gqa_ratio = num_qo_heads // num_kv_heads
    
    # Allocate outputs
    output = torch.zeros((batch_size, num_qo_heads, head_dim), dtype=torch.bfloat16, device=device)
    lse = torch.full((batch_size, num_qo_heads), -float('inf'), dtype=torch.float32, device=device)
    
    # Flatten k_cache and v_cache for page_size=1
    # Shape: [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, head_dim]
    k_cache_flat = k_cache.squeeze(1)
    v_cache_flat = v_cache.squeeze(1)
    
    # Configure kernel - optimized for B200
    BLOCK_M = 128   # Process more tokens per block for B200
    
    # Launch kernel
    grid = (batch_size, num_qo_heads)
    
    gqa_paged_decode_kernel_optimized[grid](
        q, k_cache_flat, v_cache_flat,
        kv_indptr, kv_indices,
        output, lse,
        sm_scale,
        batch_size, num_pages,
        BLOCK_M=BLOCK_M,
        NUM_QO_HEADS=num_qo_heads,
        NUM_KV_HEADS=num_kv_heads,
        HEAD_DIM=head_dim,
        GQA_RATIO=gqa_ratio,
    )
    
    return output, lse
scrolls · 323 lines total

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

Best evidence level for this revision: reported

JSON