submission 512649
mreso · python · License unknown
Kernel source · 166 lines ↓holds 1 record
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 166 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-512649?include=source"interfacepython
Compatibility
measured onNVIDIA L4
declared hardwareNVIDIA L4
architecturessm_89
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:de1d14b5576788db58ff52293870575aef9a3e2d08a83797e0ad4ffe30e6ca88
license declaredunknown
license concludedunknown
authorsmreso
imported2026-08-15
Kernel source
submission.py166 lines
# submission.py
# Adaptive prefix sum: algorithm selected per GPU at runtime.
#
# B200 / H100 / A100 → two-pass (reduce → cumsum → scan)
# Rationale: these GPUs have high HBM bandwidth; the two-pass already
# saturates it, and its 3 simple kernel launches beat the spin-polling
# overhead of decoupled lookback.
#
# L4 → single-pass decoupled lookback (Merrill & Garland 2016)
# Rationale: L4's narrower bandwidth (300 GB/s) makes the 33% reduction
# in data movement (3N → 2N) worthwhile despite the lookback overhead.
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# ════════════════════════════════════════════════════════════════════════════
# Shared helper
# ════════════════════════════════════════════════════════════════════════════
@triton.jit
def _add(a, b):
return a + b
# ════════════════════════════════════════════════════════════════════════════
# Two-pass kernels (B200 / H100 / A100)
# ════════════════════════════════════════════════════════════════════════════
@triton.jit
def _block_reduce(in_ptr, sums_ptr, N: int, BLOCK: tl.constexpr):
"""Pass 1: reduce each block of BLOCK elements to a single sum."""
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < N
x = tl.load(in_ptr + offs, mask=mask, other=0.0)
tl.store(sums_ptr + pid, tl.sum(x, axis=0))
@triton.jit
def _final_scan(in_ptr, out_ptr, sums_ptr, N: int, BLOCK: tl.constexpr):
"""Pass 2: local inclusive scan + add inter-block prefix."""
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < N
x = tl.load(in_ptr + offs, mask=mask, other=0.0)
x_scan = tl.associative_scan(x, 0, _add)
# Prefix = inclusive cumsum of all preceding blocks (stored in sums after
# the torch.cumsum call in the launcher).
prev = tl.maximum(pid - 1, 0)
prefix_raw = tl.load(sums_ptr + prev)
prefix = tl.where(pid > 0, prefix_raw, 0.0)
tl.store(out_ptr + offs, x_scan + prefix, mask=mask)
def _two_pass(x: torch.Tensor, out: torch.Tensor, N: int,
BLOCK: int, WARPS: int) -> torch.Tensor:
P = triton.cdiv(N, BLOCK)
sums = torch.empty(P, dtype=torch.float32, device=x.device)
_block_reduce[(P,)](x, sums, N, BLOCK=BLOCK, num_warps=WARPS)
if P > 1:
torch.cumsum(sums, dim=0, out=sums)
_final_scan[(P,)](x, out, sums, N, BLOCK=BLOCK, num_warps=WARPS)
return out
# ════════════════════════════════════════════════════════════════════════════
# Single-pass decoupled lookback kernel (L4)
# ════════════════════════════════════════════════════════════════════════════
@triton.jit
def _scan_kernel(
in_ptr, out_ptr,
agg_ptr, # float32[P]: local aggregate (immutable after PARTIAL)
inc_ptr, # float32[P]: inclusive prefix (immutable after COMPLETE)
status_ptr, # int32[P]: 0=invalid, 1=partial, 2=complete
N: int,
BLOCK: tl.constexpr,
):
tile_id = tl.program_id(0)
offs = tile_id * BLOCK + tl.arange(0, BLOCK)
mask = offs < N
x = tl.load(in_ptr + offs, mask=mask, other=0.0)
x_scan = tl.associative_scan(x, 0, _add)
sel = tl.arange(0, BLOCK) == BLOCK - 1
local_agg = tl.sum(tl.where(sel, x_scan, 0.0))
# Publish PARTIAL (release ensures agg store is visible before status)
tl.store(agg_ptr + tile_id, local_agg)
tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')
# Lookback — separate agg/inc arrays avoid the TOCTOU race; '.cv' bypasses
# the non-coherent per-SM L1 cache, reading from the coherent L2.
excl = 0.0
look = tile_id - 1
while look >= 0:
st = tl.atomic_add(status_ptr + look, 0, sem='acquire')
if st == 2: # COMPLETE
excl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')
look = -1
elif st == 1: # PARTIAL
excl = excl + tl.load(agg_ptr + look, cache_modifier='.cv')
look = look - 1
# st == 0 (INVALID): retry same 'look'
# Publish COMPLETE
tl.store(inc_ptr + tile_id, excl + local_agg)
tl.atomic_xchg(status_ptr + tile_id, 2, sem='release')
tl.store(out_ptr + offs, x_scan + excl, mask=mask)
def _single_pass(x: torch.Tensor, out: torch.Tensor, N: int,
BLOCK: int, WARPS: int) -> torch.Tensor:
P = triton.cdiv(N, BLOCK)
agg = torch.zeros(P, dtype=torch.float32, device=x.device)
inc = torch.zeros(P, dtype=torch.float32, device=x.device)
status = torch.zeros(P, dtype=torch.int32, device=x.device)
_scan_kernel[(P,)](x, out, agg, inc, status, N, BLOCK=BLOCK, num_warps=WARPS)
return out
# ════════════════════════════════════════════════════════════════════════════
# GPU detection (cached)
# ════════════════════════════════════════════════════════════════════════════
_CFG: tuple | None = None # (algo, BLOCK, WARPS)
def _get_cfg():
global _CFG
if _CFG is None:
name = torch.cuda.get_device_name(0).lower()
if 'l4' in name:
_CFG = ('single', 4096, 4)
elif 'b200' in name:
_CFG = ('two', 4096, 8)
elif 'h100' in name:
_CFG = ('two', 4096, 8)
elif 'a100' in name:
_CFG = ('two', 4096, 8)
else:
_CFG = ('two', 4096, 8)
return _CFG
# ════════════════════════════════════════════════════════════════════════════
# Entry point
# ════════════════════════════════════════════════════════════════════════════
def custom_kernel(data: input_t) -> output_t:
x, out = data
N = x.numel()
if N == 0:
return out
algo, BLOCK, WARPS = _get_cfg()
if algo == 'single':
return _single_pass(x, out, N, BLOCK, WARPS)
else:
return _two_pass(x, out, N, BLOCK, WARPS)
scrolls · 166 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 512647.
# submission.py- # Single-pass inclusive prefix sum — Decoupled Lookback (Merrill & Garland 2016)+ # Adaptive prefix sum: algorithm selected per GPU at runtime.#- # Each CTA (tile):- # 1. Computes a local inclusive scan.- # 2. Publishes its LOCAL AGGREGATE to agg_ptr[] as STATUS_PARTIAL.- # 3. Looks back at preceding tiles to accumulate the exclusive prefix:- # - PARTIAL tile → add its local aggregate from agg_ptr[], keep looking.- # - COMPLETE tile → add its inclusive prefix from inc_ptr[], stop.- # 4. Publishes its INCLUSIVE PREFIX to inc_ptr[] as STATUS_COMPLETE.- # 5. Adds the exclusive prefix to the local scan and writes the output.+ # B200 / H100 / A100 → two-pass (reduce → cumsum → scan)+ # Rationale: these GPUs have high HBM bandwidth; the two-pass already+ # saturates it, and its 3 simple kernel launches beat the spin-polling+ # overhead of decoupled lookback.#- # Two separate arrays (agg_ptr / inc_ptr) are required for correctness.- # With a single shared array the TOCTOU race between reading status=PARTIAL- # and loading the prefix value can yield a stale COMPLETE value; the caller- # would then treat the full inclusive prefix as a mere local aggregate and- # continue looking back, double-counting the preceding tiles.- #- # cache_modifier='.cv' (volatile — bypass L1) on all lookback loads ensures- # we read the coherent L2 value rather than a stale per-SM L1 entry.- #- # Data movement: 1×N read + 1×N write (vs 3×N for two-pass).+ # L4 → single-pass decoupled lookback (Merrill & Garland 2016)+ # Rationale: L4's narrower bandwidth (300 GB/s) makes the 33% reduction+ # in data movement (3N → 2N) worthwhile despite the lookback overhead.import torchimport triton⋯ 1 unchanged linesfrom task import input_t, output_t+ # ════════════════════════════════════════════════════════════════════════════+ # Shared helper+ # ════════════════════════════════════════════════════════════════════════════+@triton.jitdef _add(a, b):return a + b+ # ════════════════════════════════════════════════════════════════════════════+ # Two-pass kernels (B200 / H100 / A100)+ # ════════════════════════════════════════════════════════════════════════════+@triton.jit+ def _block_reduce(in_ptr, sums_ptr, N: int, BLOCK: tl.constexpr):+ """Pass 1: reduce each block of BLOCK elements to a single sum."""+ pid = tl.program_id(0)+ offs = pid * BLOCK + tl.arange(0, BLOCK)+ mask = offs < N+ x = tl.load(in_ptr + offs, mask=mask, other=0.0)+ tl.store(sums_ptr + pid, tl.sum(x, axis=0))+++ @triton.jit+ def _final_scan(in_ptr, out_ptr, sums_ptr, N: int, BLOCK: tl.constexpr):+ """Pass 2: local inclusive scan + add inter-block prefix."""+ pid = tl.program_id(0)+ offs = pid * BLOCK + tl.arange(0, BLOCK)+ mask = offs < N+ x = tl.load(in_ptr + offs, mask=mask, other=0.0)+ x_scan = tl.associative_scan(x, 0, _add)+ # Prefix = inclusive cumsum of all preceding blocks (stored in sums after+ # the torch.cumsum call in the launcher).+ prev = tl.maximum(pid - 1, 0)+ prefix_raw = tl.load(sums_ptr + prev)+ prefix = tl.where(pid > 0, prefix_raw, 0.0)+ tl.store(out_ptr + offs, x_scan + prefix, mask=mask)+++ def _two_pass(x: torch.Tensor, out: torch.Tensor, N: int,+ BLOCK: int, WARPS: int) -> torch.Tensor:+ P = triton.cdiv(N, BLOCK)+ sums = torch.empty(P, dtype=torch.float32, device=x.device)+ _block_reduce[(P,)](x, sums, N, BLOCK=BLOCK, num_warps=WARPS)+ if P > 1:+ torch.cumsum(sums, dim=0, out=sums)+ _final_scan[(P,)](x, out, sums, N, BLOCK=BLOCK, num_warps=WARPS)+ return out+++ # ════════════════════════════════════════════════════════════════════════════+ # Single-pass decoupled lookback kernel (L4)+ # ════════════════════════════════════════════════════════════════════════════++ @triton.jitdef _scan_kernel(in_ptr, out_ptr,- agg_ptr, # float32[P]: per-tile LOCAL aggregate (set at PARTIAL, immutable after)- inc_ptr, # float32[P]: per-tile INCLUSIVE prefix (set at COMPLETE, immutable after)- status_ptr, # int32[P]: 0 = invalid, 1 = partial, 2 = complete+ agg_ptr, # float32[P]: local aggregate (immutable after PARTIAL)+ inc_ptr, # float32[P]: inclusive prefix (immutable after COMPLETE)+ status_ptr, # int32[P]: 0=invalid, 1=partial, 2=completeN: int,BLOCK: tl.constexpr,):tile_id = tl.program_id(0)- # ── 1. Load tile and compute local inclusive scan ────────────────────────offs = tile_id * BLOCK + tl.arange(0, BLOCK)mask = offs < Nx = tl.load(in_ptr + offs, mask=mask, other=0.0)x_scan = tl.associative_scan(x, 0, _add)- # Local aggregate = sum of this tile's elements.sel = tl.arange(0, BLOCK) == BLOCK - 1local_agg = tl.sum(tl.where(sel, x_scan, 0.0))- # ── 2. Publish PARTIAL ───────────────────────────────────────────────────- # Store aggregate; release-fence on the xchg flushes it to L2 before any- # reader can observe status == PARTIAL.+ # Publish PARTIAL (release ensures agg store is visible before status)tl.store(agg_ptr + tile_id, local_agg)- tl.atomic_xchg(status_ptr + tile_id, 1, sem='release') # STATUS_PARTIAL+ tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')- # ── 3. Lookback ──────────────────────────────────────────────────────────- # Single loop: poll status, only advance 'look' once something is published.- # agg_ptr[j] is immutable after PARTIAL is set → safe to read whenever st>=1.- # inc_ptr[j] is immutable after COMPLETE is set → safe to read when st==2.- # This eliminates the TOCTOU race that arises with a shared prefix array.- #- # '.cv' (volatile) on loads bypasses the non-coherent per-SM L1 cache,- # reading directly from the coherent L2.+ # Lookback — separate agg/inc arrays avoid the TOCTOU race; '.cv' bypasses+ # the non-coherent per-SM L1 cache, reading from the coherent L2.excl = 0.0look = tile_id - 1while look >= 0:st = tl.atomic_add(status_ptr + look, 0, sem='acquire')- if st == 2: # COMPLETE: inclusive prefix is ready — done+ if st == 2: # COMPLETEexcl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')- look = -1 # break- elif st == 1: # PARTIAL: only the local aggregate is ready+ look = -1+ elif st == 1: # PARTIALexcl = excl + tl.load(agg_ptr + look, cache_modifier='.cv')look = look - 1- # st == 0 (INVALID): retry the same 'look' index+ # st == 0 (INVALID): retry same 'look'- # ── 4. Publish COMPLETE ──────────────────────────────────────────────────+ # Publish COMPLETEtl.store(inc_ptr + tile_id, excl + local_agg)- tl.atomic_xchg(status_ptr + tile_id, 2, sem='release') # STATUS_COMPLETE+ tl.atomic_xchg(status_ptr + tile_id, 2, sem='release')- # ── 5. Write final output ────────────────────────────────────────────────tl.store(out_ptr + offs, x_scan + excl, mask=mask)- # ---------------------------------------------------------------------------- # GPU-specific (BLOCK, num_warps) config, cached after first detection- # ---------------------------------------------------------------------------- _GPU_CFG: tuple | None = None+ def _single_pass(x: torch.Tensor, out: torch.Tensor, N: int,+ BLOCK: int, WARPS: int) -> torch.Tensor:+ P = triton.cdiv(N, BLOCK)+ agg = torch.zeros(P, dtype=torch.float32, device=x.device)+ inc = torch.zeros(P, dtype=torch.float32, device=x.device)+ status = torch.zeros(P, dtype=torch.int32, device=x.device)+ _scan_kernel[(P,)](x, out, agg, inc, status, N, BLOCK=BLOCK, num_warps=WARPS)+ return out- def _get_cfg() -> tuple[int, int]:- global _GPU_CFG- if _GPU_CFG is None:+ # ════════════════════════════════════════════════════════════════════════════+ # GPU detection (cached)+ # ════════════════════════════════════════════════════════════════════════════++ _CFG: tuple | None = None # (algo, BLOCK, WARPS)+++ def _get_cfg():+ global _CFG+ if _CFG is None:name = torch.cuda.get_device_name(0).lower()- if 'b200' in name:- _GPU_CFG = (4096, 8)+ if 'l4' in name:+ _CFG = ('single', 4096, 4)+ elif 'b200' in name:+ _CFG = ('two', 4096, 8)elif 'h100' in name:- _GPU_CFG = (4096, 8)+ _CFG = ('two', 4096, 8)elif 'a100' in name:- _GPU_CFG = (4096, 8)- elif 'l4' in name:- _GPU_CFG = (4096, 4)+ _CFG = ('two', 4096, 8)else:- _GPU_CFG = (4096, 8)- return _GPU_CFG+ _CFG = ('two', 4096, 8)+ return _CFG+ # ════════════════════════════════════════════════════════════════════════════+ # Entry point+ # ════════════════════════════════════════════════════════════════════════════+def custom_kernel(data: input_t) -> output_t:x, out = dataN = x.numel()if N == 0:return out- BLOCK, WARPS = _get_cfg()- P = triton.cdiv(N, BLOCK)-- # All three state arrays must be zeroed every call.- agg = torch.zeros(P, dtype=torch.float32, device=x.device)- inc = torch.zeros(P, dtype=torch.float32, device=x.device)- status = torch.zeros(P, dtype=torch.int32, device=x.device)-- _scan_kernel[(P,)](x, out, agg, inc, status, N, BLOCK=BLOCK, num_warps=WARPS)-- return out+ algo, BLOCK, WARPS = _get_cfg()+ if algo == 'single':+ return _single_pass(x, out, N, BLOCK, WARPS)+ else:+ return _two_pass(x, out, N, BLOCK, WARPS)
scrolls · 234 diff lines total
Best evidence level for this revision: reported
JSON