gpt-5 / triton0de5b5
gpt-5_triton_0de5b5 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 118 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-0de5b5?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:58aacca9df2b462724997435b7292d8eeb032761f7016688d6d1ce6ebd2092af
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.py118 lines
import torch
import triton
import triton.language as tl
@triton.jit
def fused_add_rmsnorm_h4096_kernel(
hidden_ptr, residual_ptr, weight_ptr, output_ptr,
M, # number of rows (batch size)
stride_hs_m, stride_hs_n,
stride_res_m, stride_res_n,
stride_out_m, stride_out_n,
H: tl.constexpr, # hidden size, must be 4096
EPS: tl.constexpr, # epsilon for numerical stability
BLOCK_SIZE: tl.constexpr,
):
tl.static_assert(H == 4096)
pid = tl.program_id(0)
row_in_bounds = pid < M
cols = tl.arange(0, BLOCK_SIZE)
# First pass: compute sum of squares across the row to get RMS
sumsq = tl.zeros([1], dtype=tl.float32)
for col_start in range(0, H, BLOCK_SIZE):
off = col_start + cols
mask = row_in_bounds & (off < H)
hs = tl.load(hidden_ptr + pid * stride_hs_m + off * stride_hs_n, mask=mask, other=0).to(tl.float32)
rs = tl.load(residual_ptr + pid * stride_res_m + off * stride_res_n, mask=mask, other=0).to(tl.float32)
x = hs + rs
sumsq += tl.sum(x * x, axis=0)
mean_sq = sumsq / H
inv_rms = tl.rsqrt(mean_sq + EPS)
# Second pass: apply normalization and weight, then store
for col_start in range(0, H, BLOCK_SIZE):
off = col_start + cols
mask = row_in_bounds & (off < H)
hs = tl.load(hidden_ptr + pid * stride_hs_m + off * stride_hs_n, mask=mask, other=0).to(tl.float32)
rs = tl.load(residual_ptr + pid * stride_res_m + off * stride_res_n, mask=mask, other=0).to(tl.float32)
w = tl.load(weight_ptr + off, mask=(off < H), other=1.0).to(tl.float32)
y = (hs + rs) * inv_rms * w
tl.store(output_ptr + pid * stride_out_m + off * stride_out_n, y.to(tl.bfloat16), mask=mask)
def run(hidden_states, residual, weight):
# Validate CUDA availability
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required to run this Triton kernel, but no CUDA device is available.")
# Validate inputs
if not isinstance(hidden_states, torch.Tensor) or not isinstance(residual, torch.Tensor) or not isinstance(weight, torch.Tensor):
raise TypeError("All inputs must be torch.Tensor instances.")
if hidden_states.ndim != 2 or residual.ndim != 2 or weight.ndim != 1:
raise ValueError("hidden_states and residual must be 2D tensors; weight must be a 1D tensor.")
if hidden_states.shape != residual.shape:
raise ValueError(f"hidden_states and residual must have the same shape, got {hidden_states.shape} vs {residual.shape}.")
B, H = hidden_states.shape
if H != 4096:
raise ValueError(f"hidden_size must be 4096, got {H}.")
if weight.shape[0] != H:
raise ValueError(f"weight must have shape ({H},), got {weight.shape}.")
# Determine target CUDA device
target_device = None
for t in (hidden_states, residual, weight):
if t.is_cuda:
target_device = t.device
break
if target_device is None:
target_device = torch.device("cuda")
# Original device of the main output (align with hidden_states)
out_device = hidden_states.device
# Move to GPU and ensure dtype is bfloat16 as specified
hs_dev = hidden_states.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
res_dev = residual.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
w_dev = weight.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
# Allocate output on GPU
out_dev = torch.empty_like(hs_dev, device=target_device, dtype=torch.bfloat16)
# Handle empty batch gracefully
if B == 0:
return out_dev.to(out_device)
# Compute strides in elements
hs_stride_m, hs_stride_n = hs_dev.stride()
res_stride_m, res_stride_n = res_dev.stride()
out_stride_m, out_stride_n = out_dev.stride()
# Launch kernel
grid = (B,)
fused_add_rmsnorm_h4096_kernel[grid](
hs_dev, res_dev, w_dev, out_dev,
B,
hs_stride_m, hs_stride_n,
res_stride_m, res_stride_n,
out_stride_m, out_stride_n,
H=4096,
EPS=1e-5,
BLOCK_SIZE=256,
num_warps=8,
num_stages=2,
)
# Move result back to original device of hidden_states
if out_device.type == "cuda" and out_device != target_device:
return out_dev.to(out_device, non_blocking=True)
elif out_device.type != "cuda":
return out_dev.to(out_device)
else:
return out_devscrolls · 118 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON