gpt-o3 / triton793f87
gpt-o3_triton_793f87 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 111 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-793f87?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:9aecf55f4e835ea63c3f587303dea2cd44adbdff5cb2cba9ec795826a3443ce2
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.py111 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def _rmsnorm_kernel(
hidden_ptr, # *bf16 [batch_size, hidden_size]
weight_ptr, # *bf16 [hidden_size]
out_ptr, # *bf16 [batch_size, hidden_size]
hidden_stride, # int stride between consecutive rows of hidden_ptr/out_ptr
out_stride, # int stride between consecutive rows of out_ptr
eps: tl.constexpr, # float numerical stability term
hidden_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0) # program id == row id
hidden_row_ptr = hidden_ptr + pid * hidden_stride
out_row_ptr = out_ptr + pid * out_stride
offs = tl.arange(0, BLOCK_SIZE)
# --------------------------------------------------------------------- #
# Pass 1 : compute sum of squares -> inverse RMS
# --------------------------------------------------------------------- #
rms_acc = tl.zeros([], dtype=tl.float32)
for start in range(0, hidden_size, BLOCK_SIZE):
idx = start + offs
mask = idx < hidden_size
x_bf16 = tl.load(hidden_row_ptr + idx, mask=mask, other=0)
x_f32 = x_bf16.to(tl.float32)
rms_acc += tl.sum(x_f32 * x_f32, axis=0)
inv_rms = tl.math.rsqrt(rms_acc / hidden_size + eps)
# --------------------------------------------------------------------- #
# Pass 2 : normalize and scale
# --------------------------------------------------------------------- #
for start in range(0, hidden_size, BLOCK_SIZE):
idx = start + offs
mask = idx < hidden_size
x_bf16 = tl.load(hidden_row_ptr + idx, mask=mask, other=0)
w_bf16 = tl.load(weight_ptr + idx, mask=mask, other=0)
x = x_bf16.to(tl.float32)
w = w_bf16.to(tl.float32)
y = x * inv_rms * w
y_bf16 = y.to(tl.bfloat16)
tl.store(out_row_ptr + idx, y_bf16, mask=mask)
def run(hidden_states: torch.Tensor, weight: torch.Tensor, *, eps: float = 1.0e-5):
"""
RMSNorm (hidden_size = 4096) implemented in Triton.
Arguments
---------
hidden_states : (batch_size, 4096) bfloat16
weight : (4096,) bfloat16
eps : float, numerical stability term
"""
# --------------------------- Sanity checks -------------------------- #
if hidden_states.dim() != 2:
raise ValueError("hidden_states must be 2-D [batch, hidden_size]")
batch_size, hidden_size = hidden_states.shape
if hidden_size != 4096:
raise ValueError(f"hidden_size must be 4096, got {hidden_size}")
if weight.dim() != 1 or weight.numel() != 4096:
raise ValueError("weight must have shape (4096,)")
if hidden_states.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
raise ValueError("hidden_states and weight must be torch.bfloat16")
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run Triton kernel")
# -------------------------- Device handling ------------------------- #
orig_device = hidden_states.device
hidden_cuda = hidden_states.cuda() if not hidden_states.is_cuda else hidden_states
weight_cuda = weight.cuda() if not weight.is_cuda else weight
hidden_cuda = hidden_cuda.contiguous()
weight_cuda = weight_cuda.contiguous()
output_cuda = torch.empty_like(hidden_cuda)
hidden_stride = hidden_cuda.stride(0)
out_stride = output_cuda.stride(0)
# ----------------------------- Launch ------------------------------- #
BLOCK_SIZE = 1024
grid = (batch_size,)
_rmsnorm_kernel[grid](
hidden_cuda,
weight_cuda,
output_cuda,
hidden_stride,
out_stride,
eps,
hidden_size=4096,
BLOCK_SIZE=BLOCK_SIZE,
num_warps=8,
num_stages=4,
)
# --------------------------- Return --------------------------------- #
return output_cuda if orig_device.type == "cuda" else output_cuda.cpu()scrolls · 111 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON