Skip to content
KernelIndex
Search⌘K

submission 553895

suvasis_29047 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 161 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-fp8-quant-553895?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
9.11µs
#5 of 17
2026-03-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5de683e5843e5d1d81bd4d6e47f2cdc0f15f036406823ab1c1ad778c814ec54a
license declaredunknown
license concludedunknown
authorssuvasis_29047
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 4num_warps = 4 : 128 threads per block
stages = 2num_stages = 2: 2-stage async prefetch hides HBM latency

Kernel source

submission.py161 lines
"""
╔══════════════════════════════════════════════════════════════════════════════╗
║   FP8 Per-Group Quantization  —  Optimized Triton Submission  v3            ║
║                                                                              ║
║  Target: B200 / Blackwell  (HBM3e 8 TB/s, 50 MB L2, TMA async)             ║
║  Strategy: Hand-written @triton.jit kernel — no Helion JIT overhead         ║
╚══════════════════════════════════════════════════════════════════════════════╝

WHY HELION TIMED OUT
─────────────────────
Helion compiles to Triton on first call, then Triton compiles to PTX/SASS.
Even with static_shapes=False, this two-stage JIT takes 3–8 minutes on a
cold remote runner.  The leaderboard ranked_timeout is 420 s (7 min) which
is not enough when compilation is included in the benchmark window.

THIS APPROACH: @triton.jit directly
─────────────────────────────────────
Writing the Triton kernel directly bypasses Helion's compilation stage.
Triton's PTX compilation is still needed on first call (~30–60 s), but
the leaderboard runner pre-warms kernels before timing, so compilation
does not count against the benchmark window.

KERNEL DESIGN
─────────────
  Grid: (T_tiles * G,)  — all (token_tile, group) pairs fully parallel
  Each thread block:
    • Owns [BLOCK_T tokens, 1 group] = [BLOCK_T, group_size] elements
    • Loads x_tile from HBM
    • Computes absmax via warp reduction (stays in registers)
    • Computes scale = absmax / 448.0
    • Quantizes: q = clamp(x / scale, -448, 448)
    • Writes x_q and x_s

  BLOCK_T = 16 : 16 tokens per block; good occupancy, low register pressure
  num_warps = 4 : 128 threads per block
  num_stages = 2: 2-stage async prefetch hides HBM latency
"""

from task import input_t, output_t

import torch
import triton
import triton.language as tl

FP8_MAX = 448.0
FP8_MIN = -448.0
FP8_EPS = 1e-10

BLOCK_T = 16
NUM_WARPS = 4
NUM_STAGES = 2


# ─────────────────────────────────────────────────────────────────────────────
# Triton kernel — fused absmax + scale + quantize
# ─────────────────────────────────────────────────────────────────────────────
@triton.jit
def _fp8_quant_triton(
    x_ptr,          # (T, H) float32 input
    xq_ptr,         # (T, H) float32 output (quantized)
    xs_ptr,         # (T, G) float32 output (scales)
    T, H, G,
    group_size,
    stride_xt, stride_xh,
    stride_qt, stride_qh,
    stride_st, stride_sg,
    BLOCK_T: tl.constexpr,
    BLOCK_GS: tl.constexpr,    # = group_size, constexpr for reduction
    FP8_MAX_VAL: tl.constexpr,
    FP8_MIN_VAL: tl.constexpr,
    FP8_EPS_VAL: tl.constexpr,
):
    # Grid: pid = t_tile * G + g_idx
    pid    = tl.program_id(0)
    T_tiles = tl.cdiv(T, BLOCK_T)
    t_tile = pid // G
    g_idx  = pid  % G

    t_off = t_tile * BLOCK_T + tl.arange(0, BLOCK_T)   # [BLOCK_T]
    gs_off = tl.arange(0, BLOCK_GS)                     # [BLOCK_GS]

    t_mask = t_off < T

    # Column range for this group
    col_start = g_idx * group_size

    # ── Load x_tile: [BLOCK_T, BLOCK_GS] ────────────────────────────────────
    x_ptrs = x_ptr + t_off[:, None] * stride_xt + (col_start + gs_off[None, :]) * stride_xh
    x_tile = tl.load(x_ptrs, mask=t_mask[:, None], other=0.0).to(tl.float32)

    # ── Per-token absmax over group_size elements ─────────────────────────────
    absmax = tl.max(tl.abs(x_tile), axis=1)             # [BLOCK_T]
    absmax = tl.maximum(absmax, FP8_EPS_VAL)

    # ── Scale ─────────────────────────────────────────────────────────────────
    scale = absmax / FP8_MAX_VAL                         # [BLOCK_T]

    # ── Quantize + clamp ──────────────────────────────────────────────────────
    q = x_tile / scale[:, None]
    q = tl.minimum(tl.maximum(q, FP8_MIN_VAL), FP8_MAX_VAL)

    # ── Store x_q ─────────────────────────────────────────────────────────────
    xq_ptrs = xq_ptr + t_off[:, None] * stride_qt + (col_start + gs_off[None, :]) * stride_qh
    tl.store(xq_ptrs, q.to(tl.float32), mask=t_mask[:, None])

    # ── Store x_s ─────────────────────────────────────────────────────────────
    xs_ptrs = xs_ptr + t_off * stride_st + g_idx * stride_sg
    tl.store(xs_ptrs, scale.to(tl.float32), mask=t_mask)


# ─────────────────────────────────────────────────────────────────────────────
# Python launcher — handles variable group_size via constexpr dispatch
# ─────────────────────────────────────────────────────────────────────────────
def _launch_fp8_quant(x, x_q, x_s, group_size):
    T, H = x.shape
    G = x_s.shape[1]

    T_tiles = triton.cdiv(T, BLOCK_T)
    grid = (T_tiles * G,)

    # BLOCK_GS must be a power-of-2 constexpr >= group_size
    # All benchmark shapes use group_size=64 or 128
    BLOCK_GS = triton.next_power_of_2(group_size)

    _fp8_quant_triton[grid](
        x, x_q, x_s,
        T, H, G,
        group_size,
        x.stride(0), x.stride(1),
        x_q.stride(0), x_q.stride(1),
        x_s.stride(0), x_s.stride(1),
        BLOCK_T=BLOCK_T,
        BLOCK_GS=BLOCK_GS,
        FP8_MAX_VAL=FP8_MAX,
        FP8_MIN_VAL=FP8_MIN,
        FP8_EPS_VAL=FP8_EPS,
        num_warps=NUM_WARPS,
        num_stages=NUM_STAGES,
    )


# ─────────────────────────────────────────────────────────────────────────────
# Public entry point
# ─────────────────────────────────────────────────────────────────────────────
def custom_kernel(data: input_t) -> output_t:
    """
    data = (x, x_q, x_s)
      x   : (T, H)   float32 CUDA  — input activations
      x_q : (T, H)   float32 CUDA  — pre-allocated output (quantized values)
      x_s : (T, G)   float32 CUDA  — pre-allocated output (scales)

    Returns (x_q, x_s) filled in-place.
    """
    x, x_q, x_s = data
    H = x.shape[1]
    G = x_s.shape[1]
    group_size = H // G

    _launch_fp8_quant(x, x_q, x_s, group_size)
    return x_q, x_s
scrolls · 161 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON