claude-opus-4-1 / tritonf7dd1f
claude-opus-4-1_triton_f7dd1f · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 110 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-f7dd1f?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
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:ef136729f5a42d9a133ba0763661771f91c62d0d208d67c976c67713fec9c542
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py110 lines
import torch
import triton
import triton.language as tl
@triton.jit
def rmsnorm_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
hidden_size: tl.constexpr,
batch_size,
BLOCK_SIZE: tl.constexpr,
):
# Each program handles one row (batch element)
row_idx = tl.program_id(axis=0)
if row_idx >= batch_size:
return
# Compute row offset
row_start = row_idx * hidden_size
# First pass: compute sum of squares
sum_squares = 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
# Load hidden states
x = tl.load(hidden_states_ptr + row_start + col_offsets, mask=mask, other=0.0).to(tl.float32)
sum_squares += tl.sum(x * x, axis=0)
# Compute inverse RMS
eps = 1e-6
mean_squares = sum_squares / hidden_size
inv_rms = 1.0 / tl.sqrt(mean_squares + eps)
# Second pass: normalize and scale
for block_start in range(0, hidden_size, BLOCK_SIZE):
col_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = col_offsets < hidden_size
# Load hidden states and weights
x = tl.load(hidden_states_ptr + row_start + col_offsets, mask=mask, other=0.0).to(tl.float32)
w = tl.load(weight_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)
# Normalize and scale
y = x * inv_rms * w
# Store output
tl.store(output_ptr + row_start + col_offsets, y.to(tl.bfloat16), mask=mask)
def run(hidden_states, weight):
# Input validation
assert hidden_states.ndim == 2, f"Expected 2D tensor, got {hidden_states.ndim}D"
assert weight.ndim == 1, f"Expected 1D tensor, got {weight.ndim}D"
batch_size, hidden_size = hidden_states.shape
assert hidden_size == 1536, f"Expected hidden_size=1536, got {hidden_size}"
assert weight.shape[0] == hidden_size, f"Weight shape mismatch: {weight.shape[0]} != {hidden_size}"
# Device management
input_device = hidden_states.device
weight_device = weight.device
# Move to GPU if needed
if hidden_states.device.type == 'cpu':
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU computation 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 computation is required")
weight = weight.cuda()
# Ensure tensors are on the same device
if hidden_states.device != weight.device:
weight = weight.to(hidden_states.device)
# Convert to bfloat16 if needed
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)
# Determine block size - optimize for B200's larger register file
BLOCK_SIZE = 256 # B200 has good memory bandwidth, use larger blocks
# Launch kernel with one thread block per batch element
grid = (batch_size,)
rmsnorm_kernel[grid](
hidden_states,
weight,
output,
hidden_size,
batch_size,
BLOCK_SIZE,
)
# Move output back to original device if needed
if input_device.type == 'cpu':
output = output.cpu()
return outputscrolls · 110 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON