Skip to content
KernelIndex
Search⌘K

submission 512667

mreso · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-512667?include=source"
interfacepython
Compatibility
measured onNVIDIA A100
declared hardwareNVIDIA A100
architecturessm_80
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
Inclusive prefix sumsuite of 11 cases
NVIDIA A100
1.81ms
#11 of 25
2026-03-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:238970ece299f0ae8a0082d9bb11f957085884b42c5bb97ff2ed2220740b8918
license declaredunknown
license concludedunknown
authorsmreso
imported2026-08-15

Techniques

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

persistent-kerneldef _persistent_scan_kernel(

Kernel source

submission.py408 lines
# submission.py
# Adaptive prefix sum: algorithm selected per GPU at runtime.
#
#   B200                →  two-pass  (already #1, near 100% BW efficiency)
#
#   H100 / A100         →  persistent single-pass decoupled lookback
#       Non-persistent single-pass launches one CTA per tile; tiles with large
#       IDs spin-wait in chains proportional to P=65536, which serialises the
#       GPU and wastes ~2x the theoretical time.
#       Persistent variant: launch exactly SM_count CTAs.  Each CTA claims
#       tiles via an atomic counter and processes them in sequence.  This
#       guarantees O(1) average lookback depth (each tile's predecessor is
#       already being computed on the same wave) and keeps all SMs 100% busy.
#
#   L4                  →  non-persistent single-pass decoupled lookback
#       L4 has fewer SMs and lower bandwidth; the persistent overhead from
#       the atomic counter is proportionally larger.  The non-persistent
#       variant is already #1 on L4.

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


# ════════════════════════════════════════════════════════════════════════════
#  Two-pass kernels  (B200)
# ════════════════════════════════════════════════════════════════════════════

@triton.jit
def _block_reduce(in_ptr, sums_ptr, N: int, BLOCK: tl.constexpr):
    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):
    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)
    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, out, N, BLOCK, WARPS):
    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


# ════════════════════════════════════════════════════════════════════════════
#  Persistent single-pass decoupled lookback  (H100 / A100)
#
#  Each CTA loops: claim a tile via atomic counter → local scan → publish
#  PARTIAL → lookback → publish COMPLETE → write output → repeat.
#
#  Because the counter assigns tile 0 before tile 1 before tile 2 …, the
#  "wave" of tiles currently in flight is always a contiguous window of
#  width ≤ num_ctas.  Tile i's predecessor (i-1) is being processed by
#  another CTA in the same wave and will publish PARTIAL/COMPLETE in at
#  most one tile-scan latency (~1 µs), bounding spin time tightly.
#
#  Separate agg_ptr / inc_ptr arrays eliminate the TOCTOU race.
#  cache_modifier='.cv' on lookback loads bypasses the non-coherent L1.
# ════════════════════════════════════════════════════════════════════════════

@triton.jit
def _persistent_scan_kernel(
    in_ptr, out_ptr,
    agg_ptr,      # float32[P]: local aggregate  (set at PARTIAL, immutable after)
    inc_ptr,      # float32[P]: inclusive prefix (set at COMPLETE, immutable after)
    status_ptr,   # int32[P]:   0=invalid, 1=partial, 2=complete
    counter_ptr,  # int32[1]:   atomic tile counter
    N: int, P: int,
    BLOCK: tl.constexpr,
):
    # Claim first tile
    tile_id = tl.atomic_add(counter_ptr, 1)

    while tile_id < P:
        # ── 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)

        sel       = tl.arange(0, BLOCK) == BLOCK - 1
        local_agg = tl.sum(tl.where(sel, x_scan, 0.0))

        # ── publish PARTIAL ──────────────────────────────────────────────
        tl.store(agg_ptr + tile_id, local_agg)
        tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')

        # ── lookback ─────────────────────────────────────────────────────
        excl = 0.0
        look = tile_id - 1
        while look >= 0:
            st = tl.atomic_add(status_ptr + look, 0, sem='acquire')
            if st == 2:
                excl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')
                look = -1
            elif st == 1:
                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')

        # ── write output ─────────────────────────────────────────────────
        tl.store(out_ptr + offs, x_scan + excl, mask=mask)

        # Claim next tile
        tile_id = tl.atomic_add(counter_ptr, 1)


def _persistent_single_pass(x, out, N, BLOCK, WARPS, num_ctas):
    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)
    counter = torch.zeros(1, dtype=torch.int32,   device=x.device)
    _persistent_scan_kernel[(num_ctas,)](
        x, out, agg, inc, status, counter, N, P,
        BLOCK=BLOCK, num_warps=WARPS,
    )
    return out


# ════════════════════════════════════════════════════════════════════════════
#  Non-persistent single-pass  (L4)
# ════════════════════════════════════════════════════════════════════════════

@triton.jit
def _nonpersistent_scan_kernel(
    in_ptr, out_ptr,
    agg_ptr, inc_ptr, status_ptr,
    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))

    tl.store(agg_ptr + tile_id, local_agg)
    tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')

    excl = 0.0
    look = tile_id - 1
    while look >= 0:
        st = tl.atomic_add(status_ptr + look, 0, sem='acquire')
        if st == 2:
            excl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')
            look = -1
        elif st == 1:
            excl = excl + tl.load(agg_ptr + look, cache_modifier='.cv')
            look = look - 1

    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 _nonpersistent_single_pass(x, out, N, BLOCK, WARPS):
    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)
    _nonpersistent_scan_kernel[(P,)](
        x, out, agg, inc, status, N, BLOCK=BLOCK, num_warps=WARPS,
    )
    return out


# ════════════════════════════════════════════════════════════════════════════
#  GPU detection (cached)
# ════════════════════════════════════════════════════════════════════════════

_CFG = None   # (algo, BLOCK, WARPS, extra)


def _get_cfg():
    global _CFG
    if _CFG is None:
        props = torch.cuda.get_device_properties(0)
        name  = props.name.lower()
        nsm   = props.multi_processor_count
        if 'l4' in name:
            _CFG = ('nonpersistent', 4096, 4, None)
        elif 'b200' in name:
            _CFG = ('two',           4096, 8, None)
        elif 'h100' in name:
            # Persistent: launch 2× SM count so each SM runs 2 concurrent
            # CTAs — fills the warp slots (64 warps/SM ÷ 8 warps/CTA = 8
            # slots; 2 CTAs × 8 warps = 16 warps, leaving headroom for the
            # other 6 persistent CTAs to spin without blocking progress).
            _CFG = ('persistent',    4096, 8, nsm * 2)
        elif 'a100' in name:
            _CFG = ('persistent',    4096, 8, nsm * 2)
        else:
            _CFG = ('two',           4096, 8, None)
    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, extra = _get_cfg()
    if algo == 'two':
        return _two_pass(x, out, N, BLOCK, WARPS)
    elif algo == 'persistent':
        return _persistent_single_pass(x, out, N, BLOCK, WARPS, num_ctas=extra)
    else:
        return _nonpersistent_single_pass(x, out, N, BLOCK, WARPS)

#
#   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 · 408 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 512651.

# submission.py
# Adaptive prefix sum: algorithm selected per GPU at runtime.
#
+ # B200 → two-pass (already #1, near 100% BW efficiency)
+ #
+ # H100 / A100 → persistent single-pass decoupled lookback
+ # Non-persistent single-pass launches one CTA per tile; tiles with large
+ # IDs spin-wait in chains proportional to P=65536, which serialises the
+ # GPU and wastes ~2x the theoretical time.
+ # Persistent variant: launch exactly SM_count CTAs. Each CTA claims
+ # tiles via an atomic counter and processes them in sequence. This
+ # guarantees O(1) average lookback depth (each tile's predecessor is
+ # already being computed on the same wave) and keeps all SMs 100% busy.
+ #
+ # L4 → non-persistent single-pass decoupled lookback
+ # L4 has fewer SMs and lower bandwidth; the persistent overhead from
+ # the atomic counter is proportionally larger. The non-persistent
+ # variant is already #1 on L4.
+
+ 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
+
+
+ # ════════════════════════════════════════════════════════════════════════════
+ # Two-pass kernels (B200)
+ # ════════════════════════════════════════════════════════════════════════════
+
+ @triton.jit
+ def _block_reduce(in_ptr, sums_ptr, N: int, BLOCK: tl.constexpr):
+ 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):
+ 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)
+ 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, out, N, BLOCK, WARPS):
+ 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
+
+
+ # ════════════════════════════════════════════════════════════════════════════
+ # Persistent single-pass decoupled lookback (H100 / A100)
+ #
+ # Each CTA loops: claim a tile via atomic counter → local scan → publish
+ # PARTIAL → lookback → publish COMPLETE → write output → repeat.
+ #
+ # Because the counter assigns tile 0 before tile 1 before tile 2 …, the
+ # "wave" of tiles currently in flight is always a contiguous window of
+ # width ≤ num_ctas. Tile i's predecessor (i-1) is being processed by
+ # another CTA in the same wave and will publish PARTIAL/COMPLETE in at
+ # most one tile-scan latency (~1 µs), bounding spin time tightly.
+ #
+ # Separate agg_ptr / inc_ptr arrays eliminate the TOCTOU race.
+ # cache_modifier='.cv' on lookback loads bypasses the non-coherent L1.
+ # ════════════════════════════════════════════════════════════════════════════
+
+ @triton.jit
+ def _persistent_scan_kernel(
+ in_ptr, out_ptr,
+ agg_ptr, # float32[P]: local aggregate (set at PARTIAL, immutable after)
+ inc_ptr, # float32[P]: inclusive prefix (set at COMPLETE, immutable after)
+ status_ptr, # int32[P]: 0=invalid, 1=partial, 2=complete
+ counter_ptr, # int32[1]: atomic tile counter
+ N: int, P: int,
+ BLOCK: tl.constexpr,
+ ):
+ # Claim first tile
+ tile_id = tl.atomic_add(counter_ptr, 1)
+
+ while tile_id < P:
+ # ── 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)
+
+ sel = tl.arange(0, BLOCK) == BLOCK - 1
+ local_agg = tl.sum(tl.where(sel, x_scan, 0.0))
+
+ # ── publish PARTIAL ──────────────────────────────────────────────
+ tl.store(agg_ptr + tile_id, local_agg)
+ tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')
+
+ # ── lookback ─────────────────────────────────────────────────────
+ excl = 0.0
+ look = tile_id - 1
+ while look >= 0:
+ st = tl.atomic_add(status_ptr + look, 0, sem='acquire')
+ if st == 2:
+ excl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')
+ look = -1
+ elif st == 1:
+ 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')
+
+ # ── write output ─────────────────────────────────────────────────
+ tl.store(out_ptr + offs, x_scan + excl, mask=mask)
+
+ # Claim next tile
+ tile_id = tl.atomic_add(counter_ptr, 1)
+
+
+ def _persistent_single_pass(x, out, N, BLOCK, WARPS, num_ctas):
+ 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)
+ counter = torch.zeros(1, dtype=torch.int32, device=x.device)
+ _persistent_scan_kernel[(num_ctas,)](
+ x, out, agg, inc, status, counter, N, P,
+ BLOCK=BLOCK, num_warps=WARPS,
+ )
+ return out
+
+
+ # ════════════════════════════════════════════════════════════════════════════
+ # Non-persistent single-pass (L4)
+ # ════════════════════════════════════════════════════════════════════════════
+
+ @triton.jit
+ def _nonpersistent_scan_kernel(
+ in_ptr, out_ptr,
+ agg_ptr, inc_ptr, status_ptr,
+ 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))
+
+ tl.store(agg_ptr + tile_id, local_agg)
+ tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')
+
+ excl = 0.0
+ look = tile_id - 1
+ while look >= 0:
+ st = tl.atomic_add(status_ptr + look, 0, sem='acquire')
+ if st == 2:
+ excl = excl + tl.load(inc_ptr + look, cache_modifier='.cv')
+ look = -1
+ elif st == 1:
+ excl = excl + tl.load(agg_ptr + look, cache_modifier='.cv')
+ look = look - 1
+
+ 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 _nonpersistent_single_pass(x, out, N, BLOCK, WARPS):
+ 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)
+ _nonpersistent_scan_kernel[(P,)](
+ x, out, agg, inc, status, N, BLOCK=BLOCK, num_warps=WARPS,
+ )
+ return out
+
+
+ # ════════════════════════════════════════════════════════════════════════════
+ # GPU detection (cached)
+ # ════════════════════════════════════════════════════════════════════════════
+
+ _CFG = None # (algo, BLOCK, WARPS, extra)
+
+
+ def _get_cfg():
+ global _CFG
+ if _CFG is None:
+ props = torch.cuda.get_device_properties(0)
+ name = props.name.lower()
+ nsm = props.multi_processor_count
+ if 'l4' in name:
+ _CFG = ('nonpersistent', 4096, 4, None)
+ elif 'b200' in name:
+ _CFG = ('two', 4096, 8, None)
+ elif 'h100' in name:
+ # Persistent: launch 2× SM count so each SM runs 2 concurrent
+ # CTAs — fills the warp slots (64 warps/SM ÷ 8 warps/CTA = 8
+ # slots; 2 CTAs × 8 warps = 16 warps, leaving headroom for the
+ # other 6 persistent CTAs to spin without blocking progress).
+ _CFG = ('persistent', 4096, 8, nsm * 2)
+ elif 'a100' in name:
+ _CFG = ('persistent', 4096, 8, nsm * 2)
+ else:
+ _CFG = ('two', 4096, 8, None)
+ 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, extra = _get_cfg()
+ if algo == 'two':
+ return _two_pass(x, out, N, BLOCK, WARPS)
+ elif algo == 'persistent':
+ return _persistent_single_pass(x, out, N, BLOCK, WARPS, num_ctas=extra)
+ else:
+ return _nonpersistent_single_pass(x, out, N, BLOCK, WARPS)
+
+ #
# 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
scrolls · 248 diff lines total

Best evidence level for this revision: reported

JSON