claude-opus-4-1 / tritonfe43bf
claude-opus-4-1_triton_fe43bf · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 120 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-fe43bf?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:dcff8838ac76c3d0184d23dd775f18f02ec311a4dd486e6a15764b33391698d1
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py120 lines
import torch
import triton
import triton.language as tl
@triton.jit
def fused_add_rmsnorm_h7168_kernel(
hidden_states_ptr,
residual_ptr,
weight_ptr,
output_ptr,
batch_size,
HIDDEN_SIZE: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
EPS: tl.constexpr,
):
# Get program id for batch dimension
pid = tl.program_id(0)
if pid >= batch_size:
return
# Base pointers for this batch element
hidden_states_row = hidden_states_ptr + pid * HIDDEN_SIZE
residual_row = residual_ptr + pid * HIDDEN_SIZE
output_row = output_ptr + pid * HIDDEN_SIZE
# Accumulator for computing mean of squares
acc = 0.0
# First pass: compute sum of squares after addition
for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
mask = offset + tl.arange(0, BLOCK_SIZE) < HIDDEN_SIZE
# Load hidden_states and residual
hidden_states = tl.load(hidden_states_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
residual = tl.load(residual_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
# Add and square
x = hidden_states + residual
x_squared = x * x
# Accumulate sum
acc += tl.sum(x_squared, axis=0)
# Compute RMS normalization factor
mean = acc / HIDDEN_SIZE
inv_rms = tl.rsqrt(mean + EPS)
# Second pass: apply normalization and weight
for offset in range(0, HIDDEN_SIZE, BLOCK_SIZE):
mask = offset + tl.arange(0, BLOCK_SIZE) < HIDDEN_SIZE
# Load hidden_states and residual again
hidden_states = tl.load(hidden_states_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
residual = tl.load(residual_row + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
# Load weight
weight = tl.load(weight_ptr + offset + tl.arange(0, BLOCK_SIZE), mask=mask, other=0.0).to(tl.float32)
# Compute normalized output
x = hidden_states + residual
y = (x * inv_rms) * weight
# Store result
tl.store(output_row + offset + tl.arange(0, BLOCK_SIZE), y.to(tl.bfloat16), mask=mask)
def run(hidden_states, residual, weight):
# Validate input shapes and constants
batch_size, hidden_size = hidden_states.shape
assert hidden_size == 7168, f"Expected hidden_size=7168, 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"
# Check dtypes
assert hidden_states.dtype == torch.bfloat16, f"Expected bfloat16 for hidden_states"
assert residual.dtype == torch.bfloat16, f"Expected bfloat16 for residual"
assert weight.dtype == torch.bfloat16, f"Expected bfloat16 for weight"
# Store original device
original_device = hidden_states.device
# Move to GPU if needed
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available. This kernel requires a GPU.")
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()
# Allocate output tensor
output = torch.empty_like(hidden_states)
# Constants
HIDDEN_SIZE = 7168
BLOCK_SIZE = 256 # Optimized for B200 memory hierarchy
EPS = 1e-6
# Launch kernel with 1D grid (one thread block per batch element)
grid = (batch_size,)
fused_add_rmsnorm_h7168_kernel[grid](
hidden_states,
residual,
weight,
output,
batch_size,
HIDDEN_SIZE=HIDDEN_SIZE,
BLOCK_SIZE=BLOCK_SIZE,
EPS=EPS,
)
# Move output back to original device if needed
if original_device.type != 'cuda':
output = output.to(original_device)
return outputscrolls · 120 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON