Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / triton49f398

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-49f398?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:5605298dbdf6e589eba611717f5c9f15f6ec01ade1e524a980f6fd3905aacb5b
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.

mmaqk = tl.dot(q_block, tl.trans(k_block))
online-softmaxm_i_new = tl.maximum(m_i, m_ij)
tile-m = 64BLOCK_M = 64
tile-n = 64BLOCK_N = 64

Kernel source

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

@triton.jit
def gqa_ragged_prefill_causal_kernel(
    q_ptr, k_ptr, v_ptr,
    qo_indptr_ptr, kv_indptr_ptr,
    output_ptr, lse_ptr,
    sm_scale,
    batch_idx,
    total_q, total_kv,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    NUM_QO_HEADS: tl.constexpr,
    NUM_KV_HEADS: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    GQA_RATIO: tl.constexpr,
):
    # Get 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)
    
    num_q_tokens = q_end - q_start
    num_kv_tokens = kv_end - kv_start
    
    if num_q_tokens <= 0 or num_kv_tokens <= 0:
        return
    
    # Get the query block and head index
    pid_m = tl.program_id(0)
    pid_h = tl.program_id(1)
    
    q_block_start = pid_m * BLOCK_M
    if q_block_start >= num_q_tokens:
        return
    
    # Determine KV head for this query head (GQA)
    kv_head = pid_h // GQA_RATIO
    
    # Initialize offsets for dimensions
    offs_m = q_block_start + tl.arange(0, BLOCK_M)
    offs_d = tl.arange(0, HEAD_DIM)
    
    # Mask for valid query positions
    mask_m = offs_m < num_q_tokens
    
    # Load query block
    global_q_indices = q_start + offs_m
    q_ptrs = q_ptr + (global_q_indices[:, None] * NUM_QO_HEADS * HEAD_DIM + 
                      pid_h * HEAD_DIM + offs_d[None, :])
    q_mask = mask_m[:, None]
    q_block = tl.load(q_ptrs, mask=q_mask, other=0.0).to(tl.float32)
    
    # Initialize accumulators
    m_i = tl.full([BLOCK_M], value=-float('inf'), dtype=tl.float32)
    l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
    
    delta = num_kv_tokens - num_q_tokens
    
    # Process KV blocks
    for kv_block_start in range(0, num_kv_tokens, BLOCK_N):
        # Create KV indices
        offs_n = kv_block_start + tl.arange(0, BLOCK_N)
        kv_mask = offs_n < num_kv_tokens
        
        # Apply causal mask - compute max valid KV index for each query
        # For query at position q_idx, can attend to KV at positions 0 to (q_idx + delta)
        max_kv_idx = offs_m[:, None] + delta + 1
        causal_mask = offs_n[None, :] < max_kv_idx
        
        # Combined mask
        combined_mask = causal_mask & mask_m[:, None] & kv_mask[None, :]
        
        # Load K block
        global_kv_indices = kv_start + offs_n
        k_ptrs = k_ptr + (global_kv_indices[:, None] * NUM_KV_HEADS * HEAD_DIM + 
                         kv_head * HEAD_DIM + offs_d[None, :])
        k_mask = kv_mask[:, None]
        k_block = tl.load(k_ptrs, mask=k_mask, other=0.0).to(tl.float32)
        
        # Compute QK^T
        qk = tl.dot(q_block, tl.trans(k_block))
        
        # Apply scaling
        qk = qk * sm_scale
        
        # Apply combined mask
        qk = tl.where(combined_mask, qk, -float('inf'))
        
        # Online softmax update
        m_ij = tl.max(qk, axis=1)
        m_ij = tl.where(mask_m, m_ij, -float('inf'))
        
        # Update max values
        m_i_new = tl.maximum(m_i, m_ij)
        
        # Compute exponentials with stability
        alpha = tl.exp(m_i - m_i_new)
        p = tl.exp(qk - m_i_new[:, None])
        
        # Mask out invalid positions in p
        p = tl.where(combined_mask, p, 0.0)
        
        # Scale accumulator
        acc = acc * alpha[:, None]
        
        # Load V block
        v_ptrs = v_ptr + (global_kv_indices[:, None] * NUM_KV_HEADS * HEAD_DIM + 
                         kv_head * HEAD_DIM + offs_d[None, :])
        v_block = tl.load(v_ptrs, mask=k_mask, other=0.0).to(tl.float32)
        
        # Update accumulator
        acc += tl.dot(p, v_block)
        
        # Update sum of exponentials
        l_ij = tl.sum(p, axis=1)
        l_i = l_i * alpha + l_ij
        m_i = m_i_new
    
    # Normalize output
    l_i_safe = tl.where(l_i > 0, l_i, 1.0)
    acc = acc / l_i_safe[:, None]
    
    # Store output
    out_ptrs = output_ptr + (global_q_indices[:, None] * NUM_QO_HEADS * HEAD_DIM + 
                             pid_h * HEAD_DIM + offs_d[None, :])
    tl.store(out_ptrs, acc.to(tl.bfloat16), mask=q_mask)
    
    # Store LSE (convert to base-2 log)
    lse_ptrs = lse_ptr + global_q_indices * NUM_QO_HEADS + pid_h
    log2_e = 1.0 / math.log(2.0)
    lse_vals = tl.where(l_i > 0, (m_i + tl.log(l_i)) * log2_e, -float('inf'))
    tl.store(lse_ptrs, lse_vals, mask=mask_m)

def run(q, k, v, qo_indptr, kv_indptr, sm_scale):
    # Store original device
    original_device = q.device
    
    # Handle device management
    if not q.is_cuda:
        if torch.cuda.is_available():
            q = q.cuda()
            k = k.cuda()
            v = v.cuda()
            qo_indptr = qo_indptr.cuda()
            kv_indptr = kv_indptr.cuda()
        else:
            raise RuntimeError("CUDA is not available but GPU tensors are required")
    
    # Get dimensions
    total_q, num_qo_heads, head_dim = q.shape
    total_kv, num_kv_heads, _ = k.shape
    len_indptr = qo_indptr.shape[0]
    batch_size = len_indptr - 1
    
    # Verify constants
    assert num_qo_heads == 32
    assert num_kv_heads == 8
    assert head_dim == 128
    
    # Verify constraints
    assert total_q == qo_indptr[-1].item()
    assert total_kv == kv_indptr[-1].item()
    
    # Allocate output tensors
    output = torch.zeros((total_q, num_qo_heads, head_dim), dtype=torch.bfloat16, device=q.device)
    lse = torch.full((total_q, num_qo_heads), -float('inf'), dtype=torch.float32, device=q.device)
    
    # Kernel configuration optimized for B200
    BLOCK_M = 64
    BLOCK_N = 64
    GQA_RATIO = num_qo_heads // num_kv_heads
    
    # Launch kernel for each batch
    for batch_idx in range(batch_size):
        q_start = qo_indptr[batch_idx].item()
        q_end = qo_indptr[batch_idx + 1].item()
        num_q_tokens = q_end - q_start
        
        if num_q_tokens <= 0:
            continue
        
        grid = (triton.cdiv(num_q_tokens, BLOCK_M), num_qo_heads)
        
        gqa_ragged_prefill_causal_kernel[grid](
            q, k, v,
            qo_indptr, kv_indptr,
            output, lse,
            sm_scale,
            batch_idx,
            total_q, total_kv,
            BLOCK_M=BLOCK_M,
            BLOCK_N=BLOCK_N,
            NUM_QO_HEADS=num_qo_heads,
            NUM_KV_HEADS=num_kv_heads,
            HEAD_DIM=head_dim,
            GQA_RATIO=GQA_RATIO,
        )
    
    # Move results back to original device if necessary
    if not original_device.type == 'cuda':
        output = output.cpu()
        lse = lse.cpu()
    
    return output, lse
scrolls · 210 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON