Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / tritond347a1

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

8 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RMSNorm h512bf16 · [512] · batch_size=1
NVIDIA B200
7.75µs
#5 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=18
NVIDIA B200
8.08µs
#5 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=64
NVIDIA B200
8.09µs
#5 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=32
NVIDIA B200
8.15µs
#5 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=7
NVIDIA B200
8.15µs
#5 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=539
NVIDIA B200
8.17µs
#5 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=11949
NVIDIA B200
15.5µs
#3 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=14521
NVIDIA B200
16.4µs
#3 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:ff658c9ba7434c2c0def752683e8226f4dd818ee37dd2ef005dfc17c8d504efa
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Kernel source

main.py123 lines
import torch
import triton
import triton.language as tl

@triton.jit
def rmsnorm_h512_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    HIDDEN_SIZE: tl.constexpr,
    EPS: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    # Process one row per program
    row_idx = tl.program_id(0)
    
    if row_idx >= batch_size:
        return
    
    # Compute row offset
    row_start = row_idx * HIDDEN_SIZE
    
    # Load and accumulate in multiple passes for better precision
    acc = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
    
    for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
        cols = offset + tl.arange(0, BLOCK_SIZE)
        mask = cols < HIDDEN_SIZE
        
        # Load hidden states for this block
        hidden_states_ptrs = hidden_states_ptr + row_start + cols
        hidden_vals = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
        
        # Accumulate squared values
        acc += hidden_vals * hidden_vals
    
    # Compute mean of squared values
    mean_sq = tl.sum(acc) / HIDDEN_SIZE
    
    # Compute inverse RMS
    inv_rms = tl.rsqrt(mean_sq + EPS)
    
    # Apply normalization and weight in blocks
    for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
        cols = offset + tl.arange(0, BLOCK_SIZE)
        mask = cols < HIDDEN_SIZE
        
        # Load hidden states
        hidden_states_ptrs = hidden_states_ptr + row_start + cols
        hidden_vals = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
        
        # Load weights
        weight_ptrs = weight_ptr + cols
        weight_vals = tl.load(weight_ptrs, mask=mask, other=0.0).to(tl.float32)
        
        # Apply RMSNorm: (x * inv_rms) * weight
        normed = hidden_vals * inv_rms * weight_vals
        
        # Store output
        output_ptrs = output_ptr + row_start + cols
        tl.store(output_ptrs, normed.to(tl.bfloat16), mask=mask)


def run(hidden_states, weight):
    # Validate inputs
    assert hidden_states.ndim == 2, f"hidden_states must be 2D, got {hidden_states.ndim}D"
    assert weight.ndim == 1, f"weight must be 1D, got {weight.ndim}D"
    
    batch_size, hidden_size = hidden_states.shape
    assert hidden_size == 512, f"hidden_size must be 512, got {hidden_size}"
    assert weight.shape[0] == hidden_size, f"weight shape mismatch: {weight.shape[0]} != {hidden_size}"
    
    # Handle device management
    original_device = hidden_states.device
    
    # Move to GPU if needed
    if hidden_states.device.type == 'cpu':
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU tensors are required for Triton kernels")
        hidden_states = hidden_states.cuda()
    
    if weight.device.type == 'cpu':
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU tensors are required for Triton kernels")
        weight = weight.cuda()
    
    # Ensure both tensors are on the same GPU
    if hidden_states.device != weight.device:
        weight = weight.to(hidden_states.device)
    
    # Ensure correct dtype
    if hidden_states.dtype != torch.bfloat16:
        hidden_states = hidden_states.to(torch.bfloat16)
    if weight.dtype != torch.bfloat16:
        weight = weight.to(torch.bfloat16)
    
    # Allocate output tensor
    output = torch.empty_like(hidden_states, dtype=torch.bfloat16, device=hidden_states.device)
    
    # Constants
    HIDDEN_SIZE = 512
    EPS = 1e-6
    BLOCK_SIZE = 128  # Optimized for B200's memory hierarchy
    
    # Launch kernel with one program per row
    grid = (batch_size,)
    
    rmsnorm_h512_kernel[grid](
        hidden_states,
        weight,
        output,
        batch_size,
        HIDDEN_SIZE=HIDDEN_SIZE,
        EPS=EPS,
        BLOCK_SIZE=BLOCK_SIZE,
    )
    
    # Move result back to original device if needed
    if original_device.type == 'cpu':
        output = output.cpu()
    
    return output
scrolls · 123 lines total

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

Best evidence level for this revision: reported

JSON