Skip to content
KernelIndex
Search⌘K

submission 512647

mreso · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-512647?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
#2 of 11
2026-03-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1b8ffe055b5e351bf9f9f17c266ebe6a6c3d53bd2826f3dc217ac94e736db001
license declaredunknown
license concludedunknown
authorsmreso
imported2026-08-15

Kernel source

submission.py130 lines
# submission.py
# Single-pass inclusive prefix sum — Decoupled Lookback (Merrill & Garland 2016)
#
# 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.
#
# 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).

import torch
import triton
import triton.language as tl
from task import input_t, output_t


@triton.jit
def _add(a, b):
    return a + b


@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
    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.
    tl.store(agg_ptr + tile_id, local_agg)
    tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')  # STATUS_PARTIAL

    # ── 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.
    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
            excl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')
            look = -1  # break
        elif st == 1:  # PARTIAL: only the local aggregate is ready
            excl = excl + tl.load(agg_ptr + look, cache_modifier='.cv')
            look = look - 1
        # st == 0 (INVALID): retry the same 'look' index

    # ── 4. Publish COMPLETE ──────────────────────────────────────────────────
    tl.store(inc_ptr + tile_id, excl + local_agg)
    tl.atomic_xchg(status_ptr + tile_id, 2, sem='release')  # STATUS_COMPLETE

    # ── 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 _get_cfg() -> tuple[int, int]:
    global _GPU_CFG
    if _GPU_CFG is None:
        name = torch.cuda.get_device_name(0).lower()
        if 'b200' in name:
            _GPU_CFG = (4096, 8)
        elif 'h100' in name:
            _GPU_CFG = (4096, 8)
        elif 'a100' in name:
            _GPU_CFG = (4096, 8)
        elif 'l4' in name:
            _GPU_CFG = (4096, 4)
        else:
            _GPU_CFG = (4096, 8)
    return _GPU_CFG


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
scrolls · 130 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 512630.

- # prefixsum_v2 submission — two-pass inclusive prefix sum using Triton
+ # submission.py
+ # Single-pass inclusive prefix sum — Decoupled Lookback (Merrill & Garland 2016)
#
- # Algorithm:
- # Pass 1 (_block_reduce): each CTA sums its BLOCK elements → P block sums.
- # Intermediate: prefix-scan the P block sums with torch.cumsum.
- # Pass 2 (_final_scan): each CTA reloads its chunk, does a local inclusive
- # scan with tl.associative_scan, adds the inter-block prefix, stores.
+ # 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.
#
- # Data movement: 2×N reads + 1×N write — better than the 4×N of the naive
- # store-partial-buffer approach.
+ # 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.
#
- # Target: NVIDIA B200 (CC 10.0), Triton 3.6.0
+ # 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).
import torch
import triton
⋯ 7 unchanged lines
@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))
+ 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
+ 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)
- @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)
+ # 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))
- # Parallel inclusive prefix scan within this block
- x_scan = tl.associative_scan(x, 0, _add)
+ # ── 2. Publish PARTIAL ───────────────────────────────────────────────────
+ # Store aggregate; release-fence on the xchg flushes it to L2 before any
+ # reader can observe status == PARTIAL.
+ tl.store(agg_ptr + tile_id, local_agg)
+ tl.atomic_xchg(status_ptr + tile_id, 1, sem='release') # STATUS_PARTIAL
- # Inter-block prefix: sum of all elements in preceding blocks.
- # sums[i] holds the inclusive cumsum of block sums after the
- # torch.cumsum step, so sums[pid-1] = sum of blocks 0 .. pid-1.
- prev = tl.maximum(pid - 1, 0)
- prefix_raw = tl.load(sums_ptr + prev)
- prefix = tl.where(pid > 0, prefix_raw, 0.0)
+ # ── 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.
+ 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
+ excl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')
+ look = -1 # break
+ elif st == 1: # PARTIAL: only the local aggregate is ready
+ excl = excl + tl.load(agg_ptr + look, cache_modifier='.cv')
+ look = look - 1
+ # st == 0 (INVALID): retry the same 'look' index
- tl.store(out_ptr + offs, x_scan + prefix, mask=mask)
+ # ── 4. Publish COMPLETE ──────────────────────────────────────────────────
+ tl.store(inc_ptr + tile_id, excl + local_agg)
+ tl.atomic_xchg(status_ptr + tile_id, 2, sem='release') # STATUS_COMPLETE
+ # ── 5. Write final output ────────────────────────────────────────────────
+ tl.store(out_ptr + offs, x_scan + excl, mask=mask)
- # Block size must be a power of 2 (required by tl.associative_scan).
- # 4096 elements × 4 B = 16 KB SMEM per block — well within B200's 256 KB.
- _BLOCK = 4096
- _WARPS = 8 # 8×32 = 256 threads → 16 elements per thread
+ # ---------------------------------------------------------------------------
+ # GPU-specific (BLOCK, num_warps) config, cached after first detection
+ # ---------------------------------------------------------------------------
+ _GPU_CFG: tuple | None = None
+
+ def _get_cfg() -> tuple[int, int]:
+ global _GPU_CFG
+ if _GPU_CFG is None:
+ name = torch.cuda.get_device_name(0).lower()
+ if 'b200' in name:
+ _GPU_CFG = (4096, 8)
+ elif 'h100' in name:
+ _GPU_CFG = (4096, 8)
+ elif 'a100' in name:
+ _GPU_CFG = (4096, 8)
+ elif 'l4' in name:
+ _GPU_CFG = (4096, 4)
+ else:
+ _GPU_CFG = (4096, 8)
+ return _GPU_CFG
+
+
def custom_kernel(data: input_t) -> output_t:
x, out = data
N = x.numel()
if N == 0:
return out
- P = triton.cdiv(N, _BLOCK)
- sums = torch.empty(P, dtype=torch.float32, device=x.device)
+ BLOCK, WARPS = _get_cfg()
+ P = triton.cdiv(N, BLOCK)
- # Pass 1: per-block reductions
- _block_reduce[(P,)](x, sums, N, BLOCK=_BLOCK, num_warps=_WARPS)
+ # 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)
- # Prefix-scan the P block sums (tiny array — torch.cumsum is fine)
- if P > 1:
- torch.cumsum(sums, dim=0, out=sums)
+ _scan_kernel[(P,)](x, out, agg, inc, status, N, BLOCK=BLOCK, num_warps=WARPS)
- # Pass 2: local scan + inter-block prefix → final output
- _final_scan[(P,)](x, out, sums, N, BLOCK=_BLOCK, num_warps=_WARPS)
-
return out
scrolls · 169 diff lines total

Best evidence level for this revision: reported

JSON