Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / triton2acb3a

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

14 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RMSNorm h4096bf16 · [4096] · batch_size=64
NVIDIA B200
14.0µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=79
NVIDIA B200
14.2µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=170
NVIDIA B200
14.3µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=63
NVIDIA B200
14.3µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=34
NVIDIA B200
14.3µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=15
NVIDIA B200
14.3µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=16
NVIDIA B200
14.3µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=7
NVIDIA B200
14.3µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=1
NVIDIA B200
14.5µs
#6 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=8804
NVIDIA B200
46.5µs
#4 of 6
2025-10-16
Show all 14 measurements ›
RMSNorm h4096bf16 · [4096] · batch_size=10827
NVIDIA B200
55.2µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=11832
NVIDIA B200
58.1µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=14418
NVIDIA B200
69.7µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=14509
NVIDIA B200
70.0µs
#4 of 6
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:1afdd0a3d4ed8e3387e9bb950e11f8befdfc47014487d52b033752e250c112e2
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Kernel source

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

@triton.jit
def rmsnorm_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    hidden_size,
    batch_size,
    eps: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    # Get the row index for this program
    row_idx = tl.program_id(0)
    
    if row_idx >= batch_size:
        return
    
    # Compute row offset
    row_start = row_idx * hidden_size
    
    # First pass: compute sum of squares
    sum_squares = 0.0
    
    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
        x = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
        
        # Accumulate sum of squares
        sum_squares += tl.sum(x * x)
    
    # Compute inverse RMS
    mean_squares = sum_squares / hidden_size
    inv_rms = tl.rsqrt(mean_squares + eps)
    
    # Second pass: apply normalization and scaling
    for offset in range(0, hidden_size, BLOCK_SIZE):
        cols = offset + tl.arange(0, BLOCK_SIZE)
        mask = cols < hidden_size
        
        # Load hidden states and weights
        hidden_states_ptrs = hidden_states_ptr + row_start + cols
        weight_ptrs = weight_ptr + cols
        
        x = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
        w = tl.load(weight_ptrs, mask=mask, other=0.0).to(tl.float32)
        
        # Apply RMSNorm
        y = x * inv_rms * w
        
        # Store output
        output_ptrs = output_ptr + row_start + cols
        tl.store(output_ptrs, y.to(tl.bfloat16), mask=mask)


@triton.jit
def rmsnorm_kernel_fused(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    hidden_size,
    batch_size,
    eps: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    # Get the row and block indices
    row_idx = tl.program_id(0)
    block_idx = tl.program_id(1)
    
    if row_idx >= batch_size:
        return
    
    # Compute offsets
    row_start = row_idx * hidden_size
    block_start = block_idx * BLOCK_SIZE
    
    if block_start >= hidden_size:
        return
    
    # Load the block of hidden states
    cols = block_start + tl.arange(0, BLOCK_SIZE)
    mask = cols < hidden_size
    
    hidden_states_ptrs = hidden_states_ptr + row_start + cols
    x = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
    
    # Compute sum of squares for this block
    block_sum_squares = tl.sum(x * x)
    
    # Use atomic add to accumulate across blocks (requires inter-block sync)
    # For simplicity, we'll use the two-pass approach in the main kernel


def run(hidden_states, weight):
    # Device handling
    original_device = hidden_states.device
    
    # Move tensors to GPU if needed
    if not torch.cuda.is_available() and (hidden_states.is_cuda or weight.is_cuda):
        raise RuntimeError("CUDA is not available but GPU tensors were provided")
    
    if torch.cuda.is_available():
        if not hidden_states.is_cuda:
            hidden_states = hidden_states.cuda()
        if not weight.is_cuda:
            weight = weight.cuda()
    else:
        raise RuntimeError("CUDA is required for Triton kernel execution")
    
    batch_size, hidden_size = hidden_states.shape
    
    # Check constants
    assert hidden_size == 4096, f"hidden_size must be 4096, got {hidden_size}"
    
    # Ensure correct dtypes
    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, device=hidden_states.device)
    
    # Determine block size - optimize for B200 with larger blocks
    # B200 has high memory bandwidth, can handle larger blocks efficiently
    BLOCK_SIZE = 512  # Increased for better vectorization on B200
    
    # Launch kernel with grid
    grid = (batch_size,)
    
    eps = 1e-5
    
    rmsnorm_kernel[grid](
        hidden_states,
        weight,
        output,
        hidden_size,
        batch_size,
        eps,
        BLOCK_SIZE=BLOCK_SIZE,
    )
    
    # Move result back to original device if necessary
    if output.device != original_device:
        output = output.to(original_device)
    
    return output
scrolls · 153 lines total

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

Best evidence level for this revision: reported

JSON