claude-opus-4-1 / triton2acb3a
claude-opus-4-1_triton_2acb3a · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 153 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-2acb3a?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:1afdd0a3d4ed8e3387e9bb950e11f8befdfc47014487d52b033752e250c112e2
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py153 lines
import torch
import triton
import triton.language as tl
@triton.jit
def rmsnorm_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
hidden_size,
batch_size,
eps: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
# Get the row index for this program
row_idx = tl.program_id(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 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
x = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
# Accumulate sum of squares
sum_squares += tl.sum(x * x)
# Compute inverse RMS
mean_squares = sum_squares / hidden_size
inv_rms = tl.rsqrt(mean_squares + eps)
# Second pass: apply normalization and scaling
for offset in range(0, hidden_size, BLOCK_SIZE):
cols = offset + tl.arange(0, BLOCK_SIZE)
mask = cols < hidden_size
# Load hidden states and weights
hidden_states_ptrs = hidden_states_ptr + row_start + cols
weight_ptrs = weight_ptr + cols
x = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
w = tl.load(weight_ptrs, mask=mask, other=0.0).to(tl.float32)
# Apply RMSNorm
y = x * inv_rms * w
# Store output
output_ptrs = output_ptr + row_start + cols
tl.store(output_ptrs, y.to(tl.bfloat16), mask=mask)
@triton.jit
def rmsnorm_kernel_fused(
hidden_states_ptr,
weight_ptr,
output_ptr,
hidden_size,
batch_size,
eps: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
# Get the row and block indices
row_idx = tl.program_id(0)
block_idx = tl.program_id(1)
if row_idx >= batch_size:
return
# Compute offsets
row_start = row_idx * hidden_size
block_start = block_idx * BLOCK_SIZE
if block_start >= hidden_size:
return
# Load the block of hidden states
cols = block_start + tl.arange(0, BLOCK_SIZE)
mask = cols < hidden_size
hidden_states_ptrs = hidden_states_ptr + row_start + cols
x = tl.load(hidden_states_ptrs, mask=mask, other=0.0).to(tl.float32)
# Compute sum of squares for this block
block_sum_squares = tl.sum(x * x)
# Use atomic add to accumulate across blocks (requires inter-block sync)
# For simplicity, we'll use the two-pass approach in the main kernel
def run(hidden_states, weight):
# Device handling
original_device = hidden_states.device
# Move tensors to GPU if needed
if not torch.cuda.is_available() and (hidden_states.is_cuda or weight.is_cuda):
raise RuntimeError("CUDA is not available but GPU tensors were provided")
if torch.cuda.is_available():
if not hidden_states.is_cuda:
hidden_states = hidden_states.cuda()
if not weight.is_cuda:
weight = weight.cuda()
else:
raise RuntimeError("CUDA is required for Triton kernel execution")
batch_size, hidden_size = hidden_states.shape
# Check constants
assert hidden_size == 4096, f"hidden_size must be 4096, got {hidden_size}"
# Ensure correct dtypes
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, device=hidden_states.device)
# Determine block size - optimize for B200 with larger blocks
# B200 has high memory bandwidth, can handle larger blocks efficiently
BLOCK_SIZE = 512 # Increased for better vectorization on B200
# Launch kernel with grid
grid = (batch_size,)
eps = 1e-5
rmsnorm_kernel[grid](
hidden_states,
weight,
output,
hidden_size,
batch_size,
eps,
BLOCK_SIZE=BLOCK_SIZE,
)
# Move result back to original device if necessary
if output.device != original_device:
output = output.to(original_device)
return outputscrolls · 153 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON