Skip to content
KernelIndex
Search⌘K

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 = 2num_warps=2,
stages = 1num_stages=1

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

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON