Skip to content
KernelIndex
Search⌘K

gpt-5 / triton159afd

gpt-5_triton_159afd · gpt-5-2025-08-07 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-159afd?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 h7168bf16 · [7168] · batch_size=18
NVIDIA B200
12.3µs
#4 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=64
NVIDIA B200
12.3µs
#4 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=32
NVIDIA B200
12.3µs
#4 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=7
NVIDIA B200
12.6µs
#4 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=1
NVIDIA B200
13.2µs
#4 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=539
NVIDIA B200
14.3µs
#4 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=11949
NVIDIA B200
93.7µs
#4 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=14521
NVIDIA B200
110.8µ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:df56ce0a16c6e91a6fb710a19cc1df0daaf9742d2c0fdc53f6e2a13842c5e61c
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 = 8num_warps=8,
stages = 2num_stages=2,

Kernel source

main.py115 lines
import torch
import triton
import triton.language as tl


@triton.jit
def rmsnorm_h7168_kernel(
    x_ptr,       # *bf16 [B, N]
    w_ptr,       # *bf16 [N]
    y_ptr,       # *bf16 [B, N]
    B: tl.constexpr,
    stride_xb, stride_xn,
    stride_yb, stride_yn,
    stride_w,
    eps: tl.float32,
    BLOCK_SIZE: tl.constexpr,
    N: tl.constexpr,
):
    row = tl.program_id(0)
    offs = tl.arange(0, BLOCK_SIZE)

    # Accumulate sum of squares in fp32
    sum_sqs = tl.zeros([1], dtype=tl.float32)
    for col_start in range(0, N, BLOCK_SIZE):
        idx = col_start + offs
        mask = (row < B) & (idx < N)
        x = tl.load(x_ptr + row * stride_xb + idx * stride_xn, mask=mask, other=0).to(tl.float32)
        sum_sqs += tl.sum(x * x, axis=0)

    denom = tl.full([1], N, dtype=tl.float32)
    mean = sum_sqs / denom
    inv_rms = tl.rsqrt(mean + eps)

    # Normalize and scale by weight, write out in bf16
    for col_start in range(0, N, BLOCK_SIZE):
        idx = col_start + offs
        mask = (row < B) & (idx < N)
        x = tl.load(x_ptr + row * stride_xb + idx * stride_xn, mask=mask, other=0).to(tl.float32)
        w = tl.load(w_ptr + idx * stride_w, mask=(idx < N), other=0).to(tl.float32)
        y = (x * inv_rms) * w
        tl.store(y_ptr + row * stride_yb + idx * stride_yn, y.to(tl.bfloat16), mask=mask)


def run(hidden_states, weight):
    if not isinstance(hidden_states, torch.Tensor) or not isinstance(weight, torch.Tensor):
        raise TypeError("hidden_states and weight must be torch.Tensor")

    # Check dimensions and constants
    if hidden_states.ndim != 2:
        raise ValueError(f"hidden_states must be 2D, got shape {hidden_states.shape}")
    if weight.ndim != 1:
        raise ValueError(f"weight must be 1D, got shape {weight.shape}")

    batch_size, hidden_size = hidden_states.shape
    if hidden_size != 7168:
        raise AssertionError(f"hidden_size must be 7168, got {hidden_size}")
    if weight.shape[0] != 7168:
        raise AssertionError(f"weight must have shape [7168], got {tuple(weight.shape)}")

    if hidden_states.dtype != torch.bfloat16:
        raise TypeError(f"hidden_states must be torch.bfloat16, got {hidden_states.dtype}")
    if weight.dtype != torch.bfloat16:
        raise TypeError(f"weight must be torch.bfloat16, got {weight.dtype}")

    # Handle empty batch fast-path
    if batch_size == 0:
        return hidden_states.clone()

    # Device management
    hs_dev = hidden_states.device
    w_dev = weight.device
    cuda_available = torch.cuda.is_available()
    compute_on_cuda = cuda_available

    if not compute_on_cuda:
        # If any tensor is on CUDA but CUDA not available (shouldn't happen), or Triton is required
        if hs_dev.type == "cuda" or w_dev.type == "cuda":
            raise RuntimeError("CUDA is not available but input tensors are on CUDA.")
        # Triton requires CUDA; cannot run on CPU
        raise RuntimeError("CUDA is not available. Triton kernels require a CUDA-capable GPU.")

    # Move to CUDA if needed and ensure contiguity
    x_gpu = hidden_states.cuda() if hs_dev.type != "cuda" else hidden_states
    w_gpu = weight.cuda() if w_dev.type != "cuda" else weight
    x_gpu = x_gpu.contiguous()
    w_gpu = w_gpu.contiguous()

    # Prepare output
    y_gpu = torch.empty_like(x_gpu, device=x_gpu.device, dtype=torch.bfloat16)

    # Strides in elements
    stride_xb, stride_xn = x_gpu.stride()
    stride_yb, stride_yn = y_gpu.stride()
    stride_w = w_gpu.stride(0)

    # Launch kernel
    grid = (batch_size,)
    eps = 1e-6
    rmsnorm_h7168_kernel[grid](
        x_gpu, w_gpu, y_gpu,
        B=batch_size,
        stride_xb=stride_xb, stride_xn=stride_xn,
        stride_yb=stride_yb, stride_yn=stride_yn,
        stride_w=stride_w,
        eps=eps,
        BLOCK_SIZE=1024,
        N=7168,
        num_warps=8,
        num_stages=2,
    )

    # Move result back to original device of hidden_states
    if hs_dev.type != "cuda":
        return y_gpu.to(hs_dev)
    return y_gpu
scrolls · 115 lines total

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

Best evidence level for this revision: reported

JSON