Skip to content
KernelIndex
Search⌘K

claude-opus-4-1 / tritonb9c384

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-b9c384?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
RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
12.4µs
#5 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
12.5µs
#5 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
12.6µs
#6 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
13.8µs
#6 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
14.2µs
#6 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
149.6µs
#6 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
191.3µ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:206bfedd3f07f3c951bf49d5312e038f408f28dac838bddcbe3303e4a42ae15d
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

persistent-kernelnum_programs = tl.num_programs(0)

Kernel source

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

@triton.jit
def rmsnorm_h2048_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    hidden_size,
    BLOCK_SIZE: tl.constexpr,
):
    # Get the row index for this program
    row_idx = tl.program_id(0)
    
    # Early exit if beyond batch size
    if row_idx >= batch_size:
        return
    
    # Compute base pointer for this row
    row_offset = row_idx * hidden_size
    
    # First pass: compute sum of squares
    sum_sq = 0.0
    for block_start in range(0, hidden_size, BLOCK_SIZE):
        # Load block of values
        col_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = col_offsets < hidden_size
        
        # Load hidden states for this block
        hidden_vals = tl.load(
            hidden_states_ptr + row_offset + col_offsets,
            mask=mask,
            other=0.0
        ).to(tl.float32)
        
        # Accumulate sum of squares
        sum_sq += tl.sum(hidden_vals * hidden_vals, axis=0)
    
    # Compute inverse RMS
    eps = 1e-6
    mean_sq = sum_sq / hidden_size
    inv_rms = tl.rsqrt(mean_sq + eps)
    
    # Second pass: apply normalization and scaling
    for block_start in range(0, hidden_size, BLOCK_SIZE):
        # Load block of values
        col_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = col_offsets < hidden_size
        
        # Load hidden states and weights for this block
        hidden_vals = tl.load(
            hidden_states_ptr + row_offset + col_offsets,
            mask=mask,
            other=0.0
        ).to(tl.float32)
        
        weight_vals = tl.load(
            weight_ptr + col_offsets,
            mask=mask,
            other=0.0
        ).to(tl.float32)
        
        # Apply RMSNorm: (x * inv_rms) * weight
        normalized = hidden_vals * inv_rms
        output_vals = normalized * weight_vals
        
        # Store output
        tl.store(
            output_ptr + row_offset + col_offsets,
            output_vals.to(tl.bfloat16),
            mask=mask
        )


@triton.jit
def rmsnorm_h2048_kernel_optimized(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    hidden_size,
    BLOCK_SIZE: tl.constexpr,
):
    # This version processes multiple rows per block for better throughput
    pid = tl.program_id(0)
    num_programs = tl.num_programs(0)
    
    # Each program processes multiple rows
    rows_per_program = (batch_size + num_programs - 1) // num_programs
    row_start = pid * rows_per_program
    row_end = min(row_start + rows_per_program, batch_size)
    
    for row_idx in range(row_start, row_end):
        row_offset = row_idx * hidden_size
        
        # Compute sum of squares
        sum_sq = 0.0
        for block_start in range(0, hidden_size, BLOCK_SIZE):
            col_offsets = block_start + tl.arange(0, BLOCK_SIZE)
            mask = col_offsets < hidden_size
            
            hidden_vals = tl.load(
                hidden_states_ptr + row_offset + col_offsets,
                mask=mask,
                other=0.0
            ).to(tl.float32)
            
            sum_sq += tl.sum(hidden_vals * hidden_vals, axis=0)
        
        # Compute inverse RMS
        eps = 1e-6
        mean_sq = sum_sq / hidden_size
        inv_rms = tl.rsqrt(mean_sq + eps)
        
        # Apply normalization and scaling
        for block_start in range(0, hidden_size, BLOCK_SIZE):
            col_offsets = block_start + tl.arange(0, BLOCK_SIZE)
            mask = col_offsets < hidden_size
            
            hidden_vals = tl.load(
                hidden_states_ptr + row_offset + col_offsets,
                mask=mask,
                other=0.0
            ).to(tl.float32)
            
            weight_vals = tl.load(
                weight_ptr + col_offsets,
                mask=mask,
                other=0.0
            ).to(tl.float32)
            
            output_vals = (hidden_vals * inv_rms) * weight_vals
            
            tl.store(
                output_ptr + row_offset + col_offsets,
                output_vals.to(tl.bfloat16),
                mask=mask
            )


def run(hidden_states, weight):
    # Validate input shapes
    batch_size, hidden_size = hidden_states.shape
    assert hidden_size == 2048, f"Expected hidden_size=2048, got {hidden_size}"
    assert weight.shape == (hidden_size,), f"Expected weight shape ({hidden_size},), got {weight.shape}"
    
    # Device management
    original_device = hidden_states.device
    weight_device = weight.device
    
    # Move tensors to GPU if needed
    if hidden_states.device.type == 'cpu':
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU execution is required")
        hidden_states = hidden_states.cuda()
    
    if weight.device.type == 'cpu':
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU execution is required")
        weight = weight.cuda()
    
    # Ensure tensors are on the same device
    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)
    
    # Configure kernel launch parameters
    # B200 has good memory bandwidth, we can use larger blocks
    BLOCK_SIZE = 256  # Process 256 elements at a time
    
    # Choose kernel based on batch size
    if batch_size < 128:
        # For small batches, use one program per row
        grid = (batch_size,)
        rmsnorm_h2048_kernel[grid](
            hidden_states,
            weight,
            output,
            batch_size,
            hidden_size,
            BLOCK_SIZE=BLOCK_SIZE,
        )
    else:
        # For larger batches, use optimized kernel with work distribution
        num_sms = 132  # B200 has 132 SMs
        grid = (min(batch_size, num_sms * 4),)  # Oversubscribe for better utilization
        rmsnorm_h2048_kernel_optimized[grid](
            hidden_states,
            weight,
            output,
            batch_size,
            hidden_size,
            BLOCK_SIZE=BLOCK_SIZE,
        )
    
    # Move output back to original device if needed
    if original_device.type == 'cpu':
        output = output.cpu()
    
    return output
scrolls · 210 lines total

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

Best evidence level for this revision: reported

JSON