Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / tritonf41fa3

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-f41fa3?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
Fused add RMSNorm h4096bf16 · [4096] · batch_size=34
NVIDIA B200
20.5µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=63
NVIDIA B200
20.5µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=79
NVIDIA B200
20.6µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=170
NVIDIA B200
20.6µs
#5 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=15
NVIDIA B200
20.6µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=16
NVIDIA B200
20.7µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=64
NVIDIA B200
21.2µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=7
NVIDIA B200
21.8µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=1
NVIDIA B200
23.5µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=8804
NVIDIA B200
78.7µs
#5 of 8
2025-10-16
Show all 14 measurements ›
Fused add RMSNorm h4096bf16 · [4096] · batch_size=10827
NVIDIA B200
94.8µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=11832
NVIDIA B200
99.7µs
#5 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14418
NVIDIA B200
123.7µs
#5 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14509
NVIDIA B200
124.7µs
#5 of 8
2025-10-16

Reported · How evidence levels are derived →

Source and license

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

Kernel source

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

@triton.jit
def fused_add_rmsnorm_kernel(
    hidden_states_ptr,
    residual_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    HIDDEN_SIZE: tl.constexpr,
    EPS: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    # Get the row index for this program
    row_idx = tl.program_id(axis=0)
    
    if row_idx >= batch_size:
        return
    
    # Calculate base pointers for this row
    hidden_states_row_ptr = hidden_states_ptr + row_idx * HIDDEN_SIZE
    residual_row_ptr = residual_ptr + row_idx * HIDDEN_SIZE
    output_row_ptr = output_ptr + row_idx * HIDDEN_SIZE
    
    # Process the row in blocks
    accumulator = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
    
    # First pass: compute sum of squares for RMS normalization
    sum_sq = 0.0
    for block_start in range(0, HIDDEN_SIZE, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < HIDDEN_SIZE
        
        # Load hidden_states and residual
        hidden_states = tl.load(hidden_states_row_ptr + block_offsets, mask=mask, other=0.0).to(tl.float32)
        residual = tl.load(residual_row_ptr + block_offsets, mask=mask, other=0.0).to(tl.float32)
        
        # Add residual to hidden_states
        x = hidden_states + residual
        
        # Accumulate sum of squares
        sum_sq += tl.sum(x * x, axis=0)
    
    # Compute inverse RMS
    mean_sq = sum_sq / HIDDEN_SIZE
    inv_rms = tl.rsqrt(mean_sq + EPS)
    
    # Second pass: normalize and apply weight
    for block_start in range(0, HIDDEN_SIZE, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < HIDDEN_SIZE
        
        # Load hidden_states and residual again
        hidden_states = tl.load(hidden_states_row_ptr + block_offsets, mask=mask, other=0.0).to(tl.float32)
        residual = tl.load(residual_row_ptr + block_offsets, mask=mask, other=0.0).to(tl.float32)
        
        # Load weight
        weight = tl.load(weight_ptr + block_offsets, mask=mask, other=0.0).to(tl.float32)
        
        # Add residual, normalize, and apply weight
        x = hidden_states + residual
        y = (x * inv_rms) * weight
        
        # Store output as bfloat16
        tl.store(output_row_ptr + block_offsets, y.to(tl.bfloat16), mask=mask)


def run(hidden_states, residual, weight):
    # Check if inputs are on CPU and move to GPU if needed
    original_device = hidden_states.device
    
    if not torch.cuda.is_available():
        if hidden_states.is_cuda or residual.is_cuda or weight.is_cuda:
            raise RuntimeError("CUDA is not available but GPU tensors were provided")
        raise RuntimeError("CUDA is not available")
    
    # Move tensors to GPU if they're on CPU
    if not hidden_states.is_cuda:
        hidden_states = hidden_states.cuda()
    if not residual.is_cuda:
        residual = residual.cuda()
    if not weight.is_cuda:
        weight = weight.cuda()
    
    # Validate shapes and dtypes
    batch_size, hidden_size = hidden_states.shape
    assert hidden_size == 4096, f"Expected hidden_size=4096, got {hidden_size}"
    assert residual.shape == (batch_size, hidden_size), f"Shape mismatch for residual"
    assert weight.shape == (hidden_size,), f"Shape mismatch for weight"
    
    # Convert to bfloat16 if needed
    if hidden_states.dtype != torch.bfloat16:
        hidden_states = hidden_states.to(torch.bfloat16)
    if residual.dtype != torch.bfloat16:
        residual = residual.to(torch.bfloat16)
    if weight.dtype != torch.bfloat16:
        weight = weight.to(torch.bfloat16)
    
    # Allocate output tensor
    output = torch.empty_like(hidden_states)
    
    # Define kernel parameters
    HIDDEN_SIZE = 4096
    BLOCK_SIZE = 256  # Optimized for B200 architecture
    EPS = 1e-5
    
    # Launch kernel with one thread block per row
    grid = (batch_size,)
    
    fused_add_rmsnorm_kernel[grid](
        hidden_states,
        residual,
        weight,
        output,
        batch_size,
        HIDDEN_SIZE=HIDDEN_SIZE,
        EPS=EPS,
        BLOCK_SIZE=BLOCK_SIZE,
    )
    
    # Move output back to original device if needed
    if not original_device.type == 'cuda':
        output = output.cpu()
    
    return output
scrolls · 127 lines total

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

Best evidence level for this revision: reported

JSON