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
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-kernel
def _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