claude-opus-4-1 / tritond347a1
claude-opus-4-1_triton_d347a1 · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 123 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-d347a1?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:ff658c9ba7434c2c0def752683e8226f4dd818ee37dd2ef005dfc17c8d504efa
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py123 lines
import torch
import triton
import triton.language as tl
@triton.jit
def rmsnorm_h512_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
batch_size,
HIDDEN_SIZE: tl.constexpr,
EPS: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
# Process one row per program
row_idx = tl.program_id(0)
if row_idx >= batch_size:
return
# Compute row offset
row_start = row_idx * HIDDEN_SIZE
# Load and accumulate in multiple passes for better precision
acc = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
cols = offset + tl.arange(0, BLOCK_SIZE)
mask = cols < HIDDEN_SIZE
# Load hidden states for this block
hidden_states_ptrs = hidden_states_ptr + row_start + cols
hidden_vals = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
# Accumulate squared values
acc += hidden_vals * hidden_vals
# Compute mean of squared values
mean_sq = tl.sum(acc) / HIDDEN_SIZE
# Compute inverse RMS
inv_rms = tl.rsqrt(mean_sq + EPS)
# Apply normalization and weight in blocks
for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
cols = offset + tl.arange(0, BLOCK_SIZE)
mask = cols < HIDDEN_SIZE
# Load hidden states
hidden_states_ptrs = hidden_states_ptr + row_start + cols
hidden_vals = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
# Load weights
weight_ptrs = weight_ptr + cols
weight_vals = tl.load(weight_ptrs, mask=mask, other=0.0).to(tl.float32)
# Apply RMSNorm: (x * inv_rms) * weight
normed = hidden_vals * inv_rms * weight_vals
# Store output
output_ptrs = output_ptr + row_start + cols
tl.store(output_ptrs, normed.to(tl.bfloat16), mask=mask)
def run(hidden_states, weight):
# Validate inputs
assert hidden_states.ndim == 2, f"hidden_states must be 2D, got {hidden_states.ndim}D"
assert weight.ndim == 1, f"weight must be 1D, got {weight.ndim}D"
batch_size, hidden_size = hidden_states.shape
assert hidden_size == 512, f"hidden_size must be 512, got {hidden_size}"
assert weight.shape[0] == hidden_size, f"weight shape mismatch: {weight.shape[0]} != {hidden_size}"
# Handle device management
original_device = hidden_states.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 tensors are required for Triton kernels")
hidden_states = hidden_states.cuda()
if weight.device.type == 'cpu':
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU tensors are required for Triton kernels")
weight = weight.cuda()
# Ensure both tensors are on the same GPU
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, dtype=torch.bfloat16, device=hidden_states.device)
# Constants
HIDDEN_SIZE = 512
EPS = 1e-6
BLOCK_SIZE = 128 # Optimized for B200's memory hierarchy
# Launch kernel with one program per row
grid = (batch_size,)
rmsnorm_h512_kernel[grid](
hidden_states,
weight,
output,
batch_size,
HIDDEN_SIZE=HIDDEN_SIZE,
EPS=EPS,
BLOCK_SIZE=BLOCK_SIZE,
)
# Move result back to original device if needed
if original_device.type == 'cpu':
output = output.cpu()
return outputscrolls · 123 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON