Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / tritonc9eea2

claude-opus-4-1_triton_c9eea2 · 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-c9eea2?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16

Benchmark evidence

7 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
Fused add RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
12.4µs
#7 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
12.6µs
#7 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
13.4µs
#7 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
14.3µs
#8 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
28.2µs
#7 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
57.9µs
#6 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
70.2µs
#6 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:c97db32a3db662ee79a9e25aab2432dfb73fcca0ee685303fd38d4f2461d6e87
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 fused_add_rmsnorm_h2048_kernel(
    hidden_states_ptr,
    residual_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    HIDDEN_SIZE: tl.constexpr,
    EPS: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    # Get batch index
    batch_idx = tl.program_id(0)
    
    # Guard against out-of-bounds access
    if batch_idx >= batch_size:
        return
    
    # Compute base pointers for current batch element
    hidden_states_row = hidden_states_ptr + batch_idx * HIDDEN_SIZE
    residual_row = residual_ptr + batch_idx * HIDDEN_SIZE
    output_row = output_ptr + batch_idx * HIDDEN_SIZE
    
    # First pass: compute sum of squares for RMS normalization
    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 and add hidden_states and residual
        hidden = tl.load(hidden_states_row + cols, mask=mask, other=0.0).to(tl.float32)
        res = tl.load(residual_row + cols, mask=mask, other=0.0).to(tl.float32)
        x = hidden + res
        
        # Accumulate sum of squares
        sum_squares += tl.sum(x * x, axis=0)
    
    # Compute inverse RMS
    mean_square = sum_squares / HIDDEN_SIZE
    inv_rms = 1.0 / tl.sqrt(mean_square + EPS)
    
    # Second pass: normalize and apply weight
    for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
        cols = offset + tl.arange(0, BLOCK_SIZE)
        mask = cols < HIDDEN_SIZE
        
        # Load inputs again
        hidden = tl.load(hidden_states_row + cols, mask=mask, other=0.0).to(tl.float32)
        res = tl.load(residual_row + cols, mask=mask, other=0.0).to(tl.float32)
        x = hidden + res
        
        # Load weight
        w = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
        
        # Apply normalization and weight
        y = (x * inv_rms) * w
        
        # Store result
        tl.store(output_row + cols, y.to(tl.bfloat16), mask=mask)


def run(hidden_states, residual, weight):
    # Check if CUDA is available
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is not available. This kernel requires a GPU.")
    
    # Get original devices
    hidden_states_device = hidden_states.device
    residual_device = residual.device
    weight_device = weight.device
    
    # Move tensors to GPU if necessary
    if hidden_states.device.type != 'cuda':
        hidden_states = hidden_states.cuda()
    if residual.device.type != 'cuda':
        residual = residual.cuda()
    if weight.device.type != 'cuda':
        weight = weight.cuda()
    
    # Ensure tensors are contiguous
    hidden_states = hidden_states.contiguous()
    residual = residual.contiguous()
    weight = weight.contiguous()
    
    # Get dimensions
    batch_size, hidden_size = hidden_states.shape
    
    # Verify hidden_size
    assert hidden_size == 2048, f"Expected hidden_size=2048, got {hidden_size}"
    assert weight.shape[0] == hidden_size, f"Weight shape mismatch: {weight.shape}"
    assert residual.shape == hidden_states.shape, f"Residual shape mismatch: {residual.shape}"
    
    # Allocate output tensor
    output = torch.empty_like(hidden_states)
    
    # Define constants
    HIDDEN_SIZE = 2048
    EPS = 1e-6
    BLOCK_SIZE = 256  # Optimized for B200 architecture
    
    # Launch kernel with one thread block per batch element
    grid = (batch_size,)
    
    fused_add_rmsnorm_h2048_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 necessary
    if hidden_states_device.type != 'cuda':
        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