Skip to content
KernelIndex
Search⌘K

submission 512649

mreso · python · License unknown

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
Inclusive prefix sumsuite of 11 cases
NVIDIA L4
9.07ms
#1 of 11
2026-03-04

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 torch
import triton
⋯ 1 unchanged lines
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]: 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=complete
N: 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 < N
x = 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 - 1
local_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.0
look = tile_id - 1
while look >= 0:
st = tl.atomic_add(status_ptr + look, 0, sem='acquire')
- if st == 2: # COMPLETE: inclusive prefix is ready — done
+ if st == 2: # COMPLETE
excl = 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: # PARTIAL
excl = 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 COMPLETE
tl.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 = data
N = 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