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
Show all 14 measurements ›Showing all 14 measurements ⌄
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 outputscrolls · 127 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON