Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton19c647

gpt-o3_triton_19c647 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 112 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-19c647?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
RMSNorm h512bf16 · [512] · batch_size=539
NVIDIA B200
6.19µs
#1 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=7
NVIDIA B200
6.19µs
#1 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=1
NVIDIA B200
6.21µs
#1 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=32
NVIDIA B200
6.30µs
#2 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=64
NVIDIA B200
6.34µs
#3 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=18
NVIDIA B200
6.36µs
#3 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=11949
NVIDIA B200
18.3µs
#4 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=14521
NVIDIA B200
20.4µs
#4 of 7
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:398e112657280d9281ed964554b3ca2a1cafdca6dcd1e5ceca14979c2bb5d978
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 = 8num_warps=8,
stages = 4num_stages=4,

Kernel source

main.py112 lines
import math
import torch
import triton
import triton.language as tl

# ----------------------------------------------------------------------
# Triton kernel
# ----------------------------------------------------------------------
@triton.jit
def _rmsnorm_kernel(
    x_ptr,            # pointer to hidden_states
    w_ptr,            # pointer to weight
    o_ptr,            # pointer to output
    stride_bs,        # hidden_size stride between rows
    eps,              # epsilon for numerical stability
    BLOCK_SIZE: tl.constexpr  # hidden dimension (constant = 512)
):
    batch_id = tl.program_id(0)
    offs = tl.arange(0, BLOCK_SIZE)

    # Pointers for this row
    x_row_ptrs = x_ptr + batch_id * stride_bs + offs
    w_ptrs     = w_ptr + offs
    o_row_ptrs = o_ptr + batch_id * stride_bs + offs

    # Load inputs as bf16 then upcast to fp32 for compute
    x = tl.load(x_row_ptrs).to(tl.float32)
    w = tl.load(w_ptrs).to(tl.float32)

    # Compute inverse RMS
    mean_square = tl.sum(x * x, axis=0) / BLOCK_SIZE
    inv_rms     = tl.math.rsqrt(mean_square + eps)

    # Normalize and scale
    y = x * inv_rms * w

    # Store result (down-cast to bf16)
    tl.store(o_row_ptrs, y.to(tl.bfloat16))


# ----------------------------------------------------------------------
# Python wrapper
# ----------------------------------------------------------------------
@torch.no_grad()
def run(hidden_states: torch.Tensor,
        weight:        torch.Tensor,
        eps: float = 1e-6):
    """
    RMSNorm (hidden_size = 512) implemented with a Triton kernel.
    Args:
        hidden_states: [batch_size, 512] BF16 tensor
        weight:        [512]             BF16 tensor
        eps:           float, epsilon used for numerical stability
    Returns:
        output:        [batch_size, 512] BF16 tensor (same device as input)
    """
    # ------------------------------------------------------------------
    # Sanity checks
    # ------------------------------------------------------------------
    if hidden_states.dim() != 2:
        raise ValueError("hidden_states must be 2-D [batch_size, 512]")
    if weight.dim() != 1 or weight.numel() != 512:
        raise ValueError("weight must be 1-D with 512 elements")
    if hidden_states.shape[1] != 512:
        raise ValueError("hidden_size is fixed to 512")

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernel")

    # ------------------------------------------------------------------
    # Device management
    # ------------------------------------------------------------------
    target_device = hidden_states.device
    if target_device.type == "cpu":
        target_device = torch.device("cuda")

    hidden_states_gpu = hidden_states.to(device=target_device, dtype=torch.bfloat16, copy=False)
    weight_gpu        = weight.to(device=target_device,        dtype=torch.bfloat16, copy=False)

    # Output allocation
    output_gpu = torch.empty_like(hidden_states_gpu)

    # ------------------------------------------------------------------
    # Kernel launch configuration
    # ------------------------------------------------------------------
    batch_size  = hidden_states_gpu.shape[0]
    BLOCK_SIZE  = 512
    grid        = (batch_size,)

    _rmsnorm_kernel[grid](
        hidden_states_gpu,
        weight_gpu,
        output_gpu,
        hidden_states_gpu.stride(0),
        eps,
        BLOCK_SIZE=BLOCK_SIZE,
        num_warps=8,
        num_stages=4,
    )

    # ------------------------------------------------------------------
    # Return result on original device
    # ------------------------------------------------------------------
    if hidden_states.device.type == "cpu":
        return output_gpu.cpu()
    return output_gpu


# ----------------------------------------------------------------------
# Export for external use
# ----------------------------------------------------------------------
__all__ = ["run"]
scrolls · 112 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON