gpt-5 / tritonb6738d
gpt-5_triton_b6738d · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 116 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-b6738d?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:1cd50aed81df912a84720829a511e3757de146706f614eaa43952000be92221b
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, # tuned for B200stages = 3
num_stages=3,Kernel source
main.py116 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def rmsnorm_h1536_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
batch_size,
stride_hs_batch,
stride_hs_hidden,
stride_w,
stride_out_batch,
stride_out_hidden,
eps,
):
pid = tl.program_id(0)
if pid >= batch_size:
return
# Base pointers for the current row
hs_row_ptr = hidden_states_ptr + pid * stride_hs_batch
out_row_ptr = output_ptr + pid * stride_out_batch
# Pass 1: compute sum of squares across the hidden dimension (H=1536)
sum_sq = tl.zeros((), dtype=tl.float32)
for start in tl.static_range(0, 1536, 256):
offs = start + tl.arange(0, 256)
mask = offs < 1536
x = tl.load(hs_row_ptr + offs * stride_hs_hidden, mask=mask, other=0).to(tl.float32)
sum_sq += tl.sum(x * x, axis=0)
mean_sq = sum_sq / 1536.0
inv_rms = tl.rsqrt(mean_sq + eps)
# Pass 2: normalize and scale by weight
for start in tl.static_range(0, 1536, 256):
offs = start + tl.arange(0, 256)
mask = offs < 1536
x = tl.load(hs_row_ptr + offs * stride_hs_hidden, mask=mask, other=0).to(tl.float32)
w = tl.load(weight_ptr + offs * stride_w, mask=mask, other=0).to(tl.float32)
y = x * inv_rms * w
tl.store(out_row_ptr + offs * stride_out_hidden, y.to(tl.bfloat16), mask=mask)
def run(hidden_states, weight, eps=1e-6):
# Validate inputs
if not isinstance(hidden_states, torch.Tensor) or not isinstance(weight, torch.Tensor):
raise TypeError("hidden_states and weight must be torch.Tensor")
if hidden_states.ndim != 2:
raise ValueError(f"hidden_states must be 2D [batch_size, hidden_size], got shape {tuple(hidden_states.shape)}")
if weight.ndim != 1:
raise ValueError(f"weight must be 1D [hidden_size], got shape {tuple(weight.shape)}")
batch_size, hidden_size = hidden_states.shape
if hidden_size != 1536:
raise ValueError(f"hidden_size must be 1536, got {hidden_size}")
if weight.numel() != hidden_size:
raise ValueError(f"weight length must be {hidden_size}, got {weight.numel()}")
cuda_available = torch.cuda.is_available()
any_cuda_input = hidden_states.is_cuda or weight.is_cuda
if any_cuda_input and not cuda_available:
raise RuntimeError("CUDA is not available but GPU tensors were provided.")
if not cuda_available:
raise RuntimeError("CUDA is required to run Triton kernels, but no CUDA device is available.")
# Determine execution device
if hidden_states.is_cuda:
exec_device = hidden_states.device
elif weight.is_cuda:
exec_device = weight.device
else:
exec_device = torch.device("cuda")
# Move tensors to GPU and cast to bfloat16
hs_gpu = hidden_states.to(device=exec_device, dtype=torch.bfloat16, copy=False)
weight_gpu = weight.to(device=exec_device, dtype=torch.bfloat16, copy=False)
# Allocate output
out_gpu = torch.empty_like(hs_gpu, dtype=torch.bfloat16, device=exec_device)
# Early exit for empty batch
if batch_size == 0:
return out_gpu.to(hidden_states.device) if hidden_states.device != exec_device else out_gpu
# Strides in elements
s_hs_batch, s_hs_hidden = hs_gpu.stride()
(s_w,) = weight_gpu.stride()
s_out_batch, s_out_hidden = out_gpu.stride()
# Kernel launch configuration: one program per batch row
grid = (batch_size,)
rmsnorm_h1536_kernel[grid](
hs_gpu,
weight_gpu,
out_gpu,
batch_size,
s_hs_batch,
s_hs_hidden,
s_w,
s_out_batch,
s_out_hidden,
float(eps),
num_warps=8, # tuned for B200
num_stages=3,
)
# Move result back to original device of hidden_states
return out_gpu.to(hidden_states.device) if hidden_states.device != exec_device else out_gpuscrolls · 116 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON