Skip to content
KernelIndex
Search⌘K

claude-opus-4-1_triton_a98005

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-a98005?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:b850cb66b93e401873ac5b58f5be5ed2353fa59fdc8dc239d45266c45a3372ad
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,
stages = 1num_stages=1,

Kernel source

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

@triton.jit
def mla_paged_decode_kernel(
    q_nope_ptr, q_pe_ptr, ckv_cache_ptr, kpe_cache_ptr,
    kv_indptr_ptr, kv_indices_ptr,
    output_ptr, lse_ptr,
    sm_scale,
    batch_size,
    stride_qn_b, stride_qn_h,
    stride_qp_b, stride_qp_h,
    stride_o_b, stride_o_h,
    stride_lse_b,
    HEAD_DIM_CKV: tl.constexpr,
    HEAD_DIM_KPE: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    head_idx = tl.program_id(1)
    
    if batch_idx >= batch_size:
        return
    
    # Get KV range for this batch
    kv_start = tl.load(kv_indptr_ptr + batch_idx)
    kv_end = tl.load(kv_indptr_ptr + batch_idx + 1)
    kv_len = kv_end - kv_start
    
    if kv_len <= 0:
        # Write zeros for empty sequences
        for d_offset in range(0, HEAD_DIM_CKV, BLOCK_SIZE):
            d_range = tl.arange(0, BLOCK_SIZE) + d_offset
            mask = d_range < HEAD_DIM_CKV
            output_offset = batch_idx * stride_o_b + head_idx * stride_o_h + d_range
            tl.store(output_ptr + output_offset, tl.zeros((BLOCK_SIZE,), dtype=tl.bfloat16), mask=mask)
        
        lse_offset = batch_idx * stride_lse_b + head_idx
        tl.store(lse_ptr + lse_offset, -float('inf'))
        return
    
    # Load query vectors
    q_base = batch_idx * stride_qn_b + head_idx * stride_qn_h
    qp_base = batch_idx * stride_qp_b + head_idx * stride_qp_h
    
    # Load q_nope in chunks
    qn_chunks = []
    num_chunks = HEAD_DIM_CKV // BLOCK_SIZE
    for i in range(num_chunks):
        offset = i * BLOCK_SIZE
        d_range = tl.arange(0, BLOCK_SIZE) + offset
        qn_chunk = tl.load(q_nope_ptr + q_base + d_range).to(tl.float32)
        qn_chunks.append(qn_chunk)
    
    # Load q_pe
    qp_range = tl.arange(0, HEAD_DIM_KPE)
    qp = tl.load(q_pe_ptr + qp_base + qp_range).to(tl.float32)
    
    # Initialize accumulators
    max_logit = -float('inf')
    sum_exp = 0.0
    acc_chunks = []
    for i in range(num_chunks):
        acc_chunks.append(tl.zeros((BLOCK_SIZE,), dtype=tl.float32))
    
    # Process KV tokens in blocks for better memory efficiency
    KV_BLOCK = 32
    for kv_block_start in range(0, kv_len, KV_BLOCK):
        kv_block_end = tl.minimum(kv_block_start + KV_BLOCK, kv_len)
        
        for kv_offset in range(kv_block_end - kv_block_start):
            kv_idx = kv_block_start + kv_offset
            if kv_idx >= kv_len:
                break
                
            page_idx = tl.load(kv_indices_ptr + kv_start + kv_idx)
            kc_base = page_idx * HEAD_DIM_CKV
            kp_base = page_idx * HEAD_DIM_KPE
            
            # Compute dot product for q_nope and ckv
            dot_nope = 0.0
            kc_chunks = []
            for i in range(num_chunks):
                offset = i * BLOCK_SIZE
                d_range = tl.arange(0, BLOCK_SIZE) + offset
                kc_chunk = tl.load(ckv_cache_ptr + kc_base + d_range).to(tl.float32)
                kc_chunks.append(kc_chunk)
                dot_nope += tl.sum(qn_chunks[i] * kc_chunk)
            
            # Compute dot product for q_pe and kpe
            kp = tl.load(kpe_cache_ptr + kp_base + qp_range).to(tl.float32)
            dot_pe = tl.sum(qp * kp)
            
            # Compute scaled logit
            logit = (dot_nope + dot_pe) * sm_scale
            
            # Online softmax update
            new_max = tl.maximum(max_logit, logit)
            
            # Rescale previous accumulator
            if kv_idx > 0 and max_logit > -float('inf'):
                scale = tl.exp(max_logit - new_max)
                sum_exp *= scale
                for i in range(num_chunks):
                    acc_chunks[i] *= scale
            
            max_logit = new_max
            exp_val = tl.exp(logit - max_logit)
            sum_exp += exp_val
            
            # Accumulate weighted kc
            for i in range(num_chunks):
                acc_chunks[i] += exp_val * kc_chunks[i]
    
    # Write output
    o_base = batch_idx * stride_o_b + head_idx * stride_o_h
    if sum_exp > 0:
        inv_sum = 1.0 / sum_exp
        for i in range(num_chunks):
            offset = i * BLOCK_SIZE
            d_range = tl.arange(0, BLOCK_SIZE) + offset
            tl.store(output_ptr + o_base + d_range, 
                    (acc_chunks[i] * inv_sum).to(tl.bfloat16))
    else:
        for i in range(num_chunks):
            offset = i * BLOCK_SIZE
            d_range = tl.arange(0, BLOCK_SIZE) + offset
            tl.store(output_ptr + o_base + d_range, 
                    tl.zeros((BLOCK_SIZE,), dtype=tl.bfloat16))
    
    # Compute and store LSE (2-based)
    lse_val = -float('inf')
    if sum_exp > 0 and max_logit > -float('inf'):
        log2_e = 1.0 / math.log(2.0)
        lse_val = (max_logit + tl.log(sum_exp)) * log2_e
    
    lse_offset = batch_idx * stride_lse_b + head_idx
    tl.store(lse_ptr + lse_offset, lse_val)


def run(q_nope, q_pe, ckv_cache, kpe_cache, kv_indptr, kv_indices, sm_scale):
    # Device management
    device = q_nope.device
    original_device = device
    
    # Move to GPU if needed
    if device.type == 'cpu':
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available for GPU computation")
        device = torch.device('cuda')
        q_nope = q_nope.cuda()
        q_pe = q_pe.cuda()
        ckv_cache = ckv_cache.cuda()
        kpe_cache = kpe_cache.cuda()
        kv_indptr = kv_indptr.cuda()
        kv_indices = kv_indices.cuda()
    
    batch_size, num_qo_heads, head_dim_ckv = q_nope.shape
    head_dim_kpe = q_pe.shape[-1]
    
    # Squeeze out page_size dimension (=1)
    ckv_cache_flat = ckv_cache.squeeze(1).contiguous()
    kpe_cache_flat = kpe_cache.squeeze(1).contiguous()
    
    # Make inputs contiguous
    q_nope = q_nope.contiguous()
    q_pe = q_pe.contiguous()
    kv_indptr = kv_indptr.contiguous()
    kv_indices = kv_indices.contiguous()
    
    # Allocate outputs
    output = torch.zeros((batch_size, num_qo_heads, head_dim_ckv), 
                        dtype=torch.bfloat16, device=device)
    lse = torch.full((batch_size, num_qo_heads), -float('inf'), 
                    dtype=torch.float32, device=device)
    
    # Launch kernel with optimized configuration
    grid = (batch_size, num_qo_heads)
    
    # Use smaller block size to reduce memory usage
    BLOCK_SIZE = 64
    
    mla_paged_decode_kernel[grid](
        q_nope, q_pe, ckv_cache_flat, kpe_cache_flat,
        kv_indptr, kv_indices,
        output, lse,
        sm_scale,
        batch_size,
        q_nope.stride(0), q_nope.stride(1),
        q_pe.stride(0), q_pe.stride(1),
        output.stride(0), output.stride(1),
        lse.stride(0),
        HEAD_DIM_CKV=head_dim_ckv,
        HEAD_DIM_KPE=head_dim_kpe,
        BLOCK_SIZE=BLOCK_SIZE,
        num_warps=4,
        num_stages=1,
    )
    
    # Move back to original device if needed
    if original_device.type == 'cpu':
        output = output.cpu()
        lse = lse.cpu()
    
    return output, lse
scrolls · 207 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON