gpt-o3 / tritonc1e819
gpt-o3_triton_c1e819 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 128 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-c1e819?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:e25739502af57a6d7db7a9290eb16603f9057261e4bf11647f4e17fa984bfc5e
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
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 = 4
num_stages=4,Kernel source
main.py128 lines
import torch
import triton
import triton.language as tl
# ----------------------------------------------------------------------
# Problem-wide constants (compile-time)
# ----------------------------------------------------------------------
HIDDEN_SIZE: int = 4096 # fixed hidden dimension
BLOCK_SIZE: int = 1024 # 4 chunks per row
EPSILON: float = 1e-5 # epsilon for RMSNorm
# ----------------------------------------------------------------------
# Triton kernel
# ----------------------------------------------------------------------
@triton.jit
def _fused_add_rmsnorm_h4096(
hidden_ptr, # bf16 [B, H]
residual_ptr, # bf16 [B, H]
weight_ptr, # bf16 [H]
output_ptr, # bf16 [B, H]
stride_hidden_bs, # = HIDDEN_SIZE
stride_res_bs, # = HIDDEN_SIZE
stride_out_bs, # = HIDDEN_SIZE
BLOCK: tl.constexpr = BLOCK_SIZE,
H: tl.constexpr = HIDDEN_SIZE,
EPS: tl.constexpr = EPSILON,
):
"""
One program instance handles one row (batch element).
It iterates over the hidden dimension in BLOCK-wide chunks.
"""
# Program/id along batch dimension
pid = tl.program_id(axis=0)
# Offsets for a block of columns [0, BLOCK)
offs = tl.arange(0, BLOCK)
# Base pointers for this row
hidden_row = hidden_ptr + pid * stride_hidden_bs
residual_row = residual_ptr + pid * stride_res_bs
output_row = output_ptr + pid * stride_out_bs
# ------------------------------------------------------------------
# Pass 1: compute sum of squares
# ------------------------------------------------------------------
ssq = tl.zeros((), dtype=tl.float32)
for col in tl.static_range(0, H, BLOCK):
idx = col + offs
x_h = tl.load(hidden_row + idx).to(tl.float32)
x_r = tl.load(residual_row + idx).to(tl.float32)
x = x_h + x_r
ssq += tl.sum(x * x, axis=0)
inv_rms = tl.rsqrt(ssq / H + EPS)
# ------------------------------------------------------------------
# Pass 2: normalize, scale and store
# ------------------------------------------------------------------
for col in tl.static_range(0, H, BLOCK):
idx = col + offs
h = tl.load(hidden_row + idx).to(tl.float32)
r = tl.load(residual_row + idx).to(tl.float32)
w = tl.load(weight_ptr + idx).to(tl.float32)
y = (h + r) * inv_rms * w
tl.store(output_row + idx, y.to(tl.bfloat16))
# ----------------------------------------------------------------------
# Python wrapper
# ----------------------------------------------------------------------
def run(hidden_states: torch.Tensor,
residual: torch.Tensor,
weight: torch.Tensor,
**kwargs) -> torch.Tensor:
"""
Fused Add + RMSNorm (hidden_size = 4096, bf16) implemented in Triton.
"""
# ------------------ Sanity checks ------------------
if hidden_states.shape != residual.shape:
raise ValueError("hidden_states and residual must have identical shapes")
if hidden_states.shape[-1] != HIDDEN_SIZE:
raise ValueError(f"hidden_size must be {HIDDEN_SIZE}")
if weight.numel() != HIDDEN_SIZE:
raise ValueError(f"weight must have {HIDDEN_SIZE} elements")
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run Triton kernels")
# ------------------ Device handling ----------------
src_device = hidden_states.device
device = src_device if src_device.type == "cuda" else torch.device("cuda")
hidden_gpu = hidden_states.to(device, dtype=torch.bfloat16, non_blocking=True)
residual_gpu = residual.to(device, dtype=torch.bfloat16, non_blocking=True)
weight_gpu = weight.to(device, dtype=torch.bfloat16, non_blocking=True)
# Ensure contiguous layout
hidden_gpu = hidden_gpu.contiguous()
residual_gpu = residual_gpu.contiguous()
weight_gpu = weight_gpu.contiguous()
# ------------------ Launch kernel ------------------
batch_size = hidden_gpu.shape[0]
output_gpu = torch.empty_like(hidden_gpu)
grid = (batch_size,)
_fused_add_rmsnorm_h4096[grid](
hidden_gpu, residual_gpu, weight_gpu, output_gpu,
HIDDEN_SIZE, HIDDEN_SIZE, HIDDEN_SIZE,
num_warps=8,
num_stages=4,
)
# ------------------ Move result back ---------------
if src_device.type == "cpu":
return output_gpu.cpu()
return output_gpu
# ----------------------------------------------------------------------
# End of file
# ----------------------------------------------------------------------scrolls · 128 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON