gpt-5 / triton21bcc2
gpt-5_triton_21bcc2 · 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-21bcc2?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:fde311bab96aaa5b07c6b6d67d98fa76031e2e7a0d771525039739dc2f2f5d98
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 = 8stages = 4
num_stages = 4Kernel source
main.py118 lines
import torch
import triton
import triton.language as tl
@triton.jit
def _rmsnorm_h4096_kernel(
X_ptr, W_ptr, Y_ptr,
stride_x_row, stride_x_col,
stride_w,
stride_y_row, stride_y_col,
N: tl.constexpr,
EPSILON: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
x_row = X_ptr + pid * stride_x_row
y_row = Y_ptr + pid * stride_y_row
tl.static_assert(N == 4096)
tl.static_assert(BLOCK_SIZE > 0)
tl.static_assert(BLOCK_SIZE % 128 == 0)
sum_sq = tl.zeros((), dtype=tl.float32)
# First pass: compute sum of squares
for offs in range(0, N, BLOCK_SIZE):
idx = offs + tl.arange(0, BLOCK_SIZE)
mask = idx < N
x = tl.load(x_row + idx * stride_x_col, mask=mask, other=0).to(tl.float32)
sum_sq += tl.sum(x * x, axis=0)
mean = sum_sq / N
inv_rms = tl.rsqrt(mean + EPSILON)
# Second pass: normalize and scale
for offs in range(0, N, BLOCK_SIZE):
idx = offs + tl.arange(0, BLOCK_SIZE)
mask = idx < N
x = tl.load(x_row + idx * stride_x_col, mask=mask, other=0).to(tl.float32)
w = tl.load(W_ptr + idx * stride_w, mask=mask, other=1).to(tl.float32)
y = (x * inv_rms) * w
tl.store(y_row + idx * stride_y_col, y.to(tl.bfloat16), mask=mask)
def run(hidden_states, weight):
if hidden_states is None or weight is None:
raise ValueError("hidden_states and weight must be provided")
if hidden_states.ndim != 2:
raise ValueError(f"hidden_states must be 2D [batch, hidden], got shape {tuple(hidden_states.shape)}")
if weight.ndim != 1:
raise ValueError(f"weight must be 1D [hidden], got shape {tuple(weight.shape)}")
batch_size, hidden_size = hidden_states.shape
if hidden_size != 4096:
raise AssertionError(f"hidden_size must be 4096, got {hidden_size}")
if weight.shape[0] != hidden_size:
raise ValueError(f"weight shape mismatch: expected {hidden_size}, got {weight.shape[0]}")
# Enforce dtype
if hidden_states.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
raise TypeError("hidden_states and weight must be torch.bfloat16")
# Device management
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required to run this Triton kernel but torch.cuda.is_available() is False.")
# Select target CUDA device
target_device = None
if hidden_states.is_cuda:
target_device = hidden_states.device
elif weight.is_cuda:
target_device = weight.device
else:
target_device = torch.device('cuda')
# Move inputs to device as needed
hs_in_dev = hidden_states.to(target_device, non_blocking=True)
w_in_dev = weight.to(target_device, non_blocking=True)
# Strides in elements (PyTorch gives element strides)
stride_x_row = hs_in_dev.stride(0)
stride_x_col = hs_in_dev.stride(1)
stride_w = w_in_dev.stride(0)
stride_y_row = stride_x_row
stride_y_col = stride_x_col
# Allocate output on device
y_dev = torch.empty_like(hs_in_dev, device=target_device, dtype=torch.bfloat16)
# Launch configuration tuned for 4096 hidden size on B200
BLOCK_SIZE = 1024 # process 1024 elements per iteration, 4 iterations total
num_warps = 8
num_stages = 4
EPSILON = 1e-5
grid = (batch_size,)
_rmsnorm_h4096_kernel[grid](
hs_in_dev, w_in_dev, y_dev,
stride_x_row, stride_x_col,
stride_w,
stride_y_row, stride_y_col,
N=hidden_size,
EPSILON=EPSILON,
BLOCK_SIZE=BLOCK_SIZE,
num_warps=num_warps,
num_stages=num_stages,
)
# Move back to the original device of hidden_states
if hidden_states.device != target_device:
y_out = y_dev.to(hidden_states.device, non_blocking=True)
else:
y_out = y_dev
return y_outscrolls · 118 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON