gpt-5 / triton714ae0
gpt-5_triton_714ae0 · gpt-5-2025-08-07 · 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-gpt-5-triton-714ae0?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:767df550ec2489626c493f4196b597cddf70cccb66ee5e67524ccd39fdc0bf28
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 8
num_warps=8,stages = 2
num_stages=2,Kernel source
main.py123 lines
import torch
import triton
import triton.language as tl
@triton.jit
def fused_add_rmsnorm_h7168_kernel(
hidden_ptr, # *bfloat16
residual_ptr, # *bfloat16
weight_ptr, # *bfloat16
out_ptr, # *bfloat16
batch_size, # int32
ld_hidden, # int32
ld_residual, # int32
ld_out, # int32
eps: tl.constexpr,
H: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
row = tl.program_id(axis=0)
if row >= batch_size:
return
# Base pointers for this row
hidden_row = hidden_ptr + row * ld_hidden
residual_row = residual_ptr + row * ld_residual
out_row = out_ptr + row * ld_out
sum_sq = 0.0
# Accumulate sum of squares
for col_start in range(0, H, BLOCK_SIZE):
cols = col_start + tl.arange(0, BLOCK_SIZE)
mask = cols < H
h = tl.load(hidden_row + cols, mask=mask, other=0.0)
r = tl.load(residual_row + cols, mask=mask, other=0.0)
x = tl.cast(h, tl.float32) + tl.cast(r, tl.float32)
sum_sq += tl.sum(x * x, axis=0)
mean = sum_sq / tl.cast(H, tl.float32)
inv_rms = tl.rsqrt(mean + eps)
# Normalize, scale by weight and store
for col_start in range(0, H, BLOCK_SIZE):
cols = col_start + tl.arange(0, BLOCK_SIZE)
mask = cols < H
h = tl.load(hidden_row + cols, mask=mask, other=0.0)
r = tl.load(residual_row + cols, mask=mask, other=0.0)
w = tl.load(weight_ptr + cols, mask=mask, other=0.0)
x = tl.cast(h, tl.float32) + tl.cast(r, tl.float32)
y = x * inv_rms * tl.cast(w, tl.float32)
tl.store(out_row + cols, tl.cast(y, tl.bfloat16), mask=mask)
def run(hidden_states, residual, weight):
# Validate inputs
if not (isinstance(hidden_states, torch.Tensor) and isinstance(residual, torch.Tensor) and isinstance(weight, torch.Tensor)):
raise TypeError("All inputs must be torch.Tensor")
if hidden_states.ndim != 2 or residual.ndim != 2:
raise ValueError("hidden_states and residual must be 2D tensors of shape [batch_size, hidden_size]")
if weight.ndim != 1:
raise ValueError("weight must be 1D tensor of shape [hidden_size]")
B, H = hidden_states.shape
if H != 7168:
raise AssertionError(f"hidden_size must be 7168, got {H}")
if residual.shape != hidden_states.shape:
raise ValueError(f"residual must have same shape as hidden_states, got {residual.shape} vs {hidden_states.shape}")
if weight.shape[0] != H:
raise ValueError(f"weight must have shape [{H}], got {tuple(weight.shape)}")
# Device checks and management
devices = {hidden_states.device.type, residual.device.type, weight.device.type}
has_cuda_input = "cuda" in devices
if has_cuda_input and not torch.cuda.is_available():
raise RuntimeError("At least one input tensor is on CUDA, but CUDA is not available.")
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run the Triton kernel.")
# Determine target CUDA device
target_device = None
for t in (hidden_states, residual, weight):
if t.device.type == "cuda":
target_device = t.device
break
if target_device is None:
target_device = torch.device("cuda", torch.cuda.current_device())
# Preserve original output device
out_device = hidden_states.device
# Move to target device and ensure dtype/layout
def to_device_bf16_contig(t):
return t.to(device=target_device, dtype=torch.bfloat16, non_blocking=True).contiguous()
hs_dev = to_device_bf16_contig(hidden_states)
rs_dev = to_device_bf16_contig(residual)
w_dev = to_device_bf16_contig(weight)
# Allocate output on target device
out_dev = torch.empty_like(hs_dev, dtype=torch.bfloat16, device=target_device)
# Launch kernel
BLOCK_SIZE = 1024 # tuned for H=7168
grid = (B,)
fused_add_rmsnorm_h7168_kernel[grid](
hs_dev, rs_dev, w_dev, out_dev,
B, H, H, H,
eps=1e-6,
H=7168,
BLOCK_SIZE=BLOCK_SIZE,
num_warps=8,
num_stages=2,
)
# Move result back to original device if needed
if out_device != target_device:
out = out_dev.to(device=out_device, non_blocking=True)
else:
out = out_dev
return outscrolls · 123 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON