claude-opus-4-1 / tritonc9eea2
claude-opus-4-1_triton_c9eea2 · 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-c9eea2?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:c97db32a3db662ee79a9e25aab2432dfb73fcca0ee685303fd38d4f2461d6e87
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 fused_add_rmsnorm_h2048_kernel(
hidden_states_ptr,
residual_ptr,
weight_ptr,
output_ptr,
batch_size,
HIDDEN_SIZE: tl.constexpr,
EPS: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
# Get batch index
batch_idx = tl.program_id(0)
# Guard against out-of-bounds access
if batch_idx >= batch_size:
return
# Compute base pointers for current batch element
hidden_states_row = hidden_states_ptr + batch_idx * HIDDEN_SIZE
residual_row = residual_ptr + batch_idx * HIDDEN_SIZE
output_row = output_ptr + batch_idx * HIDDEN_SIZE
# First pass: compute sum of squares for RMS normalization
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 and add hidden_states and residual
hidden = tl.load(hidden_states_row + cols, mask=mask, other=0.0).to(tl.float32)
res = tl.load(residual_row + cols, mask=mask, other=0.0).to(tl.float32)
x = hidden + res
# Accumulate sum of squares
sum_squares += tl.sum(x * x, axis=0)
# Compute inverse RMS
mean_square = sum_squares / HIDDEN_SIZE
inv_rms = 1.0 / tl.sqrt(mean_square + EPS)
# Second pass: normalize and apply weight
for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
cols = offset + tl.arange(0, BLOCK_SIZE)
mask = cols < HIDDEN_SIZE
# Load inputs again
hidden = tl.load(hidden_states_row + cols, mask=mask, other=0.0).to(tl.float32)
res = tl.load(residual_row + cols, mask=mask, other=0.0).to(tl.float32)
x = hidden + res
# Load weight
w = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
# Apply normalization and weight
y = (x * inv_rms) * w
# Store result
tl.store(output_row + cols, y.to(tl.bfloat16), mask=mask)
def run(hidden_states, residual, weight):
# Check if CUDA is available
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available. This kernel requires a GPU.")
# Get original devices
hidden_states_device = hidden_states.device
residual_device = residual.device
weight_device = weight.device
# Move tensors to GPU if necessary
if hidden_states.device.type != 'cuda':
hidden_states = hidden_states.cuda()
if residual.device.type != 'cuda':
residual = residual.cuda()
if weight.device.type != 'cuda':
weight = weight.cuda()
# Ensure tensors are contiguous
hidden_states = hidden_states.contiguous()
residual = residual.contiguous()
weight = weight.contiguous()
# Get dimensions
batch_size, hidden_size = hidden_states.shape
# Verify hidden_size
assert hidden_size == 2048, f"Expected hidden_size=2048, got {hidden_size}"
assert weight.shape[0] == hidden_size, f"Weight shape mismatch: {weight.shape}"
assert residual.shape == hidden_states.shape, f"Residual shape mismatch: {residual.shape}"
# Allocate output tensor
output = torch.empty_like(hidden_states)
# Define constants
HIDDEN_SIZE = 2048
EPS = 1e-6
BLOCK_SIZE = 256 # Optimized for B200 architecture
# Launch kernel with one thread block per batch element
grid = (batch_size,)
fused_add_rmsnorm_h2048_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 necessary
if hidden_states_device.type != 'cuda':
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