Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / tritonfe43bf

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-fe43bf?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
Fused add RMSNorm h7168bf16 · [7168] · batch_size=64
NVIDIA B200
32.5µs
#6 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=18
NVIDIA B200
32.8µs
#6 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=32
NVIDIA B200
32.9µs
#6 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=539
NVIDIA B200
33.4µs
#6 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=7
NVIDIA B200
33.9µs
#6 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=1
NVIDIA B200
35.6µs
#6 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=11949
NVIDIA B200
191.9µs
#6 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=14521
NVIDIA B200
229.0µs
#6 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:dcff8838ac76c3d0184d23dd775f18f02ec311a4dd486e6a15764b33391698d1
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Kernel source

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

@triton.jit
def fused_add_rmsnorm_h7168_kernel(
    hidden_states_ptr,
    residual_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    HIDDEN_SIZE: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    EPS: tl.constexpr,
):
    # Get program id for batch dimension
    pid = tl.program_id(0)
    
    if pid >= batch_size:
        return
    
    # Base pointers for this batch element
    hidden_states_row = hidden_states_ptr + pid * HIDDEN_SIZE
    residual_row = residual_ptr + pid * HIDDEN_SIZE
    output_row = output_ptr + pid * HIDDEN_SIZE
    
    # Accumulator for computing mean of squares
    acc = 0.0
    
    # First pass: compute sum of squares after addition
    for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
        mask = offset + tl.arange(0, BLOCK_SIZE) < HIDDEN_SIZE
        
        # Load hidden_states and residual
        hidden_states = tl.load(hidden_states_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
        residual = tl.load(residual_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
        
        # Add and square
        x = hidden_states + residual
        x_squared = x * x
        
        # Accumulate sum
        acc += tl.sum(x_squared, axis=0)
    
    # Compute RMS normalization factor
    mean = acc / HIDDEN_SIZE
    inv_rms = tl.rsqrt(mean + EPS)
    
    # Second pass: apply normalization and weight
    for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
        mask = offset + tl.arange(0, BLOCK_SIZE) < HIDDEN_SIZE
        
        # Load hidden_states and residual again
        hidden_states = tl.load(hidden_states_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
        residual = tl.load(residual_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
        
        # Load weight
        weight = tl.load(weight_ptr + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
        
        # Compute normalized output
        x = hidden_states + residual
        y = (x * inv_rms) * weight
        
        # Store result
        tl.store(output_row + offset + tl.arange(0, BLOCK_SIZE), y.to(tl.bfloat16), mask=mask)


def run(hidden_states, residual, weight):
    # Validate input shapes and constants
    batch_size, hidden_size = hidden_states.shape
    assert hidden_size == 7168, f"Expected hidden_size=7168, 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"
    
    # Check dtypes
    assert hidden_states.dtype == torch.bfloat16, f"Expected bfloat16 for hidden_states"
    assert residual.dtype == torch.bfloat16, f"Expected bfloat16 for residual"
    assert weight.dtype == torch.bfloat16, f"Expected bfloat16 for weight"
    
    # Store original device
    original_device = hidden_states.device
    
    # Move to GPU if needed
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is not available. This kernel requires a GPU.")
    
    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()
    
    # Allocate output tensor
    output = torch.empty_like(hidden_states)
    
    # Constants
    HIDDEN_SIZE = 7168
    BLOCK_SIZE = 256  # Optimized for B200 memory hierarchy
    EPS = 1e-6
    
    # Launch kernel with 1D grid (one thread block per batch element)
    grid = (batch_size,)
    
    fused_add_rmsnorm_h7168_kernel[grid](
        hidden_states,
        residual,
        weight,
        output,
        batch_size,
        HIDDEN_SIZE=HIDDEN_SIZE,
        BLOCK_SIZE=BLOCK_SIZE,
        EPS=EPS,
    )
    
    # Move output back to original device if needed
    if original_device.type != 'cuda':
        output = output.to(original_device)
    
    return output
scrolls · 120 lines total

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

Best evidence level for this revision: reported

JSON