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
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 = 4
num_warps = 4 : 128 threads per blockstages = 2
num_stages = 2: 2-stage async prefetch hides HBM latencyKernel 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