gpt-o3 / tritond1dcce
gpt-o3_triton_d1dcce · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 89 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-d1dcce?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:ee88907127e648736813df4930948466248bfd2930e99fd68ef376b2bac5fe6c
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=4Kernel source
main.py89 lines
import torch
import triton
import triton.language as tl
EPS = 1e-6
HIDDEN_SIZE = 1536
BLOCK_SIZE = 256
@triton.jit
def rmsnorm_kernel(x_ptr, w_ptr, y_ptr,
hidden_size: tl.constexpr,
epsilon: tl.constexpr,
BLOCK_SIZE: tl.constexpr):
pid = tl.program_id(axis=0)
offs = tl.arange(0, BLOCK_SIZE)
row_start = pid * hidden_size
acc = tl.zeros((), dtype=tl.float32)
for col in tl.static_range(0, hidden_size, BLOCK_SIZE):
cols = col + offs
ptrs = x_ptr + row_start + cols
x = tl.load(ptrs, mask=cols < hidden_size, other=0)
x_f32 = x.to(tl.float32)
acc += tl.sum(x_f32 * x_f32)
mean = acc / hidden_size
inv_rms = tl.rsqrt(mean + epsilon)
for col in tl.static_range(0, hidden_size, BLOCK_SIZE):
cols = col + offs
x_ptrs = x_ptr + row_start + cols
w_ptrs = w_ptr + cols
y_ptrs = y_ptr + row_start + cols
x = tl.load(x_ptrs, mask=cols < hidden_size, other=0)
w = tl.load(w_ptrs, mask=cols < hidden_size, other=0)
y = x.to(tl.float32) * inv_rms * w.to(tl.float32)
tl.store(y_ptrs, y.to(tl.bfloat16), mask=cols < hidden_size)
def _ensure_bf16(tensor, name):
if tensor.dtype != torch.bfloat16:
raise TypeError(f"{name} must have dtype torch.bfloat16, got {tensor.dtype}")
def _to_cuda(tensor):
return tensor if tensor.is_cuda else tensor.cuda()
@torch.no_grad()
def run(*args, **kwargs):
if len(args) == 2:
hidden_states, weight = args
else:
hidden_states = kwargs.get("hidden_states")
weight = kwargs.get("weight")
if hidden_states is None or weight is None:
raise ValueError("Both 'hidden_states' and 'weight' must be provided")
_ensure_bf16(hidden_states, "hidden_states")
_ensure_bf16(weight, "weight")
if hidden_states.shape[-1] != HIDDEN_SIZE or weight.shape[0] != HIDDEN_SIZE:
raise ValueError(f"hidden_size must be {HIDDEN_SIZE}")
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available")
orig_device = hidden_states.device
x_gpu = _to_cuda(hidden_states).contiguous()
w_gpu = _to_cuda(weight).contiguous()
y_gpu = torch.empty_like(x_gpu)
grid = (x_gpu.shape[0],)
rmsnorm_kernel[grid](
x_gpu, w_gpu, y_gpu,
hidden_size=HIDDEN_SIZE,
epsilon=EPS,
BLOCK_SIZE=BLOCK_SIZE,
num_warps=8,
num_stages=4
)
return y_gpu if orig_device.type == "cuda" else y_gpu.cpu()scrolls · 89 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON