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
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-kernel
num_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 outputscrolls · 210 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON