gpt-5 / triton679e13
gpt-5_triton_679e13 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 157 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-679e13?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:93af4c1b8b13889962aa4b988589c7fecde417b3a7d9e7f7cdccf017f0796092
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.py157 lines
import torch
import triton
import triton.language as tl
# Constants
HIDDEN_SIZE = 2048
BLOCK_SIZE = 256
EPS = 1e-6
@triton.jit
def _fused_add_rmsnorm_h2048_kernel(
hidden_ptr, residual_ptr, weight_ptr, out_ptr,
B,
hs_stride_row, hs_stride_col,
rs_stride_row, rs_stride_col,
out_stride_row, out_stride_col,
w_stride,
eps,
H: tl.constexpr,
BLOCK: tl.constexpr,
):
row = tl.program_id(0)
if row >= B:
return
sumsq = 0.0
# First pass: compute sum of squares across the row
for col_start in range(0, H, BLOCK):
offs = tl.arange(0, BLOCK)
cols = col_start + offs
mask = cols < H
hs_ptrs = hidden_ptr + row * hs_stride_row + cols * hs_stride_col
rs_ptrs = residual_ptr + row * rs_stride_row + cols * rs_stride_col
hs = tl.load(hs_ptrs, mask=mask, other=0).to(tl.float32)
rs = tl.load(rs_ptrs, mask=mask, other=0).to(tl.float32)
x = hs + rs
sumsq += tl.sum(x * x, axis=0)
mean = sumsq / H
inv_rms = tl.rsqrt(mean + eps)
# Second pass: normalize and scale by weight
for col_start in range(0, H, BLOCK):
offs = tl.arange(0, BLOCK)
cols = col_start + offs
mask = cols < H
hs_ptrs = hidden_ptr + row * hs_stride_row + cols * hs_stride_col
rs_ptrs = residual_ptr + row * rs_stride_row + cols * rs_stride_col
w_ptrs = weight_ptr + cols * w_stride
out_ptrs = out_ptr + row * out_stride_row + cols * out_stride_col
hs = tl.load(hs_ptrs, mask=mask, other=0).to(tl.float32)
rs = tl.load(rs_ptrs, mask=mask, other=0).to(tl.float32)
w = tl.load(w_ptrs, mask=mask, other=0).to(tl.float32)
y = (hs + rs) * inv_rms * w
tl.store(out_ptrs, y.to(tl.bfloat16), mask=mask)
@torch.no_grad()
def run(hidden_states, residual, weight):
# 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")
if hidden_states.dtype != torch.bfloat16 or residual.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
raise TypeError("Tensors must be of dtype torch.bfloat16")
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]")
if hidden_states.shape != residual.shape:
raise ValueError("hidden_states and residual must have the same shape")
B, H = hidden_states.shape
if H != HIDDEN_SIZE:
raise ValueError(f"hidden_size must be {HIDDEN_SIZE}, got {H}")
if weight.numel() != HIDDEN_SIZE:
raise ValueError(f"weight must have {HIDDEN_SIZE} elements, got {weight.numel()}")
# Device management
cuda_available = torch.cuda.is_available()
hs_dev = hidden_states.device
rs_dev = residual.device
w_dev = weight.device
# Determine target CUDA device
target_cuda_device = None
if hs_dev.type == "cuda":
target_cuda_device = hs_dev
if rs_dev.type == "cuda":
if target_cuda_device is None:
target_cuda_device = rs_dev
elif rs_dev != target_cuda_device:
raise ValueError("All CUDA tensors must be on the same device")
if w_dev.type == "cuda":
if target_cuda_device is None:
target_cuda_device = w_dev
elif w_dev != target_cuda_device:
raise ValueError("All CUDA tensors must be on the same device")
if target_cuda_device is None:
if not cuda_available:
raise RuntimeError("CUDA is required but not available.")
target_cuda_device = torch.device("cuda")
else:
if not cuda_available:
raise RuntimeError("CUDA is not available but tensors are on CUDA.")
# Move to CUDA if needed
hs_cuda = hidden_states if hidden_states.device == target_cuda_device else hidden_states.to(device=target_cuda_device, non_blocking=True)
rs_cuda = residual if residual.device == target_cuda_device else residual.to(device=target_cuda_device, non_blocking=True)
w_cuda = weight if weight.device == target_cuda_device else weight.to(device=target_cuda_device, non_blocking=True)
# Early return for empty batch
if B == 0:
out_empty = torch.empty_like(hidden_states)
return out_empty
# Allocate output on CUDA
out_cuda = torch.empty_like(hs_cuda)
# Kernel launch
grid = (B,)
_fused_add_rmsnorm_h2048_kernel[grid](
hs_cuda,
rs_cuda,
w_cuda,
out_cuda,
B,
hs_cuda.stride(0), hs_cuda.stride(1),
rs_cuda.stride(0), rs_cuda.stride(1),
out_cuda.stride(0), out_cuda.stride(1),
w_cuda.stride(0),
float(EPS),
H=HIDDEN_SIZE,
BLOCK=BLOCK_SIZE,
num_warps=8,
num_stages=2,
)
# Move result back to original device of hidden_states
if hs_dev == out_cuda.device:
return out_cuda
else:
return out_cuda.to(hs_dev, non_blocking=True)scrolls · 157 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON