Skip to content
KernelIndex
Search⌘K

submission 523188

mreso · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

prefixsum_cub.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-523188?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.36ms
#2 of 25
2026-03-10

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

shared-memory__shared__ union {

Kernel source

prefixsum_cub.py140 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

cuda_source = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cub/cub.cuh>
#include <cub/agent/single_pass_scan_operators.cuh>

// Custom inclusive scan using CUB building blocks with optimized tile size
// and CUB's warp-cooperative decoupled lookback

constexpr int BT = 128;
constexpr int IPT = 36;
constexpr int TILE = BT * IPT;  // 4608

// cub::Sum was removed in CUDA 13 — define our own
struct SumOp {
    __host__ __device__ __forceinline__
    float operator()(const float& a, const float& b) const { return a + b; }
};

typedef cub::ScanTileState<float> TileStateT;
typedef cub::BlockLoad<float, BT, IPT, cub::BLOCK_LOAD_WARP_TRANSPOSE> BlockLoadT;
typedef cub::BlockStore<float, BT, IPT, cub::BLOCK_STORE_WARP_TRANSPOSE> BlockStoreT;
typedef cub::BlockScan<float, BT, cub::BLOCK_SCAN_WARP_SCANS> BlockScanT;
typedef cub::TilePrefixCallbackOp<float, SumOp, TileStateT> PrefixOpT;

// cub::DeviceScanInitKernel is an internal API removed in CUDA 13.
// Replicate it: call InitializeStatus on each tile state entry.
__global__ void init_tile_state_kernel(TileStateT tile_state, int num_tiles) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx < num_tiles)
        tile_state.InitializeStatus(num_tiles);
}

__global__ __launch_bounds__(BT)
void scan_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    TileStateT tile_state,
    const int n
) {
    const int tid = threadIdx.x;
    const int tile_id = blockIdx.x;
    const int tile_offset = tile_id * TILE;
    const int valid = min(TILE, n - tile_offset);

    __shared__ union {
        typename BlockLoadT::TempStorage load;
        typename BlockStoreT::TempStorage store;
        struct {
            typename BlockScanT::TempStorage scan;
            typename PrefixOpT::TempStorage prefix;
        };
    } temp;

    // Load tile data
    float items[IPT];
    if (valid == TILE)
        BlockLoadT(temp.load).Load(input + tile_offset, items);
    else
        BlockLoadT(temp.load).Load(input + tile_offset, items, valid, 0.0f);
    __syncthreads();

    // Block-level inclusive scan + inter-block lookback
    if (tile_id == 0) {
        float block_agg;
        BlockScanT(temp.scan).InclusiveSum(items, items, block_agg);
        if (tid == 0)
            tile_state.SetInclusive(0, block_agg);
    } else {
        PrefixOpT prefix_op(tile_state, temp.prefix, SumOp(), tile_id);
        BlockScanT(temp.scan).InclusiveSum(items, items, prefix_op);
    }
    __syncthreads();

    // Store results
    if (valid == TILE)
        BlockStoreT(temp.store).Store(output + tile_offset, items);
    else
        BlockStoreT(temp.store).Store(output + tile_offset, items, valid);
}

// Pre-allocated tile state storage
static void* d_tile_state_storage = nullptr;
static size_t d_tile_state_bytes = 0;
static int alloc_max_tiles = 0;

torch::Tensor inclusive_scan(torch::Tensor input, torch::Tensor output) {
    const int n = input.size(0);
    const int num_tiles = (n + TILE - 1) / TILE;

    // Grow tile state storage if needed
    if (num_tiles > alloc_max_tiles) {
        if (d_tile_state_storage) cudaFree(d_tile_state_storage);
        size_t needed;
        TileStateT::AllocationSize(num_tiles, needed);
        d_tile_state_bytes = needed * 2;
        cudaMalloc(&d_tile_state_storage, d_tile_state_bytes);
        alloc_max_tiles = num_tiles;
    }

    // Initialize tile state
    TileStateT tile_state;
    tile_state.Init(num_tiles, d_tile_state_storage, d_tile_state_bytes);
    init_tile_state_kernel<<<(num_tiles + 255) / 256, 256>>>(tile_state, num_tiles);

    // Launch scan kernel
    scan_kernel<<<num_tiles, BT>>>(
        input.data_ptr<float>(), output.data_ptr<float>(),
        tile_state, n);

    return output;
}
"""

cpp_source = r"""
torch::Tensor inclusive_scan(torch::Tensor input, torch::Tensor output);
"""

print("Compiling submission kernel...")
_module = load_inline(
    name='prefix_sum_submission_v2',
    cpp_sources=[cpp_source],
    cuda_sources=[cuda_source],
    functions=['inclusive_scan'],
    verbose=False,
    extra_cuda_cflags=['-O3', '--use_fast_math', '--threads', '4']
)
print("Done.")


def custom_kernel(data: input_t) -> output_t:
    input_tensor, output_tensor = data
    _module.inclusive_scan(input_tensor, output_tensor)
    return output_tensor
scrolls · 140 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 512667.

- # 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 torch.utils.cpp_extension import load_inline
from task import input_t, output_t
+ cuda_source = r"""
+ #include <torch/extension.h>
+ #include <cuda.h>
+ #include <cuda_runtime.h>
+ #include <cub/cub.cuh>
+ #include <cub/agent/single_pass_scan_operators.cuh>
- @triton.jit
- def _add(a, b):
- return a + b
+ // Custom inclusive scan using CUB building blocks with optimized tile size
+ // and CUB's warp-cooperative decoupled lookback
+ constexpr int BT = 128;
+ constexpr int IPT = 36;
+ constexpr int TILE = BT * IPT; // 4608
- # ════════════════════════════════════════════════════════════════════════════
- # Two-pass kernels (B200)
- # ════════════════════════════════════════════════════════════════════════════
+ // cub::Sum was removed in CUDA 13 — define our own
+ struct SumOp {
+ __host__ __device__ __forceinline__
+ float operator()(const float& a, const float& b) const { return a + b; }
+ };
- @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))
+ typedef cub::ScanTileState<float> TileStateT;
+ typedef cub::BlockLoad<float, BT, IPT, cub::BLOCK_LOAD_WARP_TRANSPOSE> BlockLoadT;
+ typedef cub::BlockStore<float, BT, IPT, cub::BLOCK_STORE_WARP_TRANSPOSE> BlockStoreT;
+ typedef cub::BlockScan<float, BT, cub::BLOCK_SCAN_WARP_SCANS> BlockScanT;
+ typedef cub::TilePrefixCallbackOp<float, SumOp, TileStateT> PrefixOpT;
+ // cub::DeviceScanInitKernel is an internal API removed in CUDA 13.
+ // Replicate it: call InitializeStatus on each tile state entry.
+ __global__ void init_tile_state_kernel(TileStateT tile_state, int num_tiles) {
+ int idx = blockIdx.x * blockDim.x + threadIdx.x;
+ if (idx < num_tiles)
+ tile_state.InitializeStatus(num_tiles);
+ }
- @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)
+ __global__ __launch_bounds__(BT)
+ void scan_kernel(
+ const float* __restrict__ input,
+ float* __restrict__ output,
+ TileStateT tile_state,
+ const int n
+ ) {
+ const int tid = threadIdx.x;
+ const int tile_id = blockIdx.x;
+ const int tile_offset = tile_id * TILE;
+ const int valid = min(TILE, n - tile_offset);
+ __shared__ union {
+ typename BlockLoadT::TempStorage load;
+ typename BlockStoreT::TempStorage store;
+ struct {
+ typename BlockScanT::TempStorage scan;
+ typename PrefixOpT::TempStorage prefix;
+ };
+ } temp;
- 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
+ // Load tile data
+ float items[IPT];
+ if (valid == TILE)
+ BlockLoadT(temp.load).Load(input + tile_offset, items);
+ else
+ BlockLoadT(temp.load).Load(input + tile_offset, items, valid, 0.0f);
+ __syncthreads();
+ // Block-level inclusive scan + inter-block lookback
+ if (tile_id == 0) {
+ float block_agg;
+ BlockScanT(temp.scan).InclusiveSum(items, items, block_agg);
+ if (tid == 0)
+ tile_state.SetInclusive(0, block_agg);
+ } else {
+ PrefixOpT prefix_op(tile_state, temp.prefix, SumOp(), tile_id);
+ BlockScanT(temp.scan).InclusiveSum(items, items, prefix_op);
+ }
+ __syncthreads();
- # ════════════════════════════════════════════════════════════════════════════
- # 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.
- # ════════════════════════════════════════════════════════════════════════════
+ // Store results
+ if (valid == TILE)
+ BlockStoreT(temp.store).Store(output + tile_offset, items);
+ else
+ BlockStoreT(temp.store).Store(output + tile_offset, items, valid);
+ }
- @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)
+ // Pre-allocated tile state storage
+ static void* d_tile_state_storage = nullptr;
+ static size_t d_tile_state_bytes = 0;
+ static int alloc_max_tiles = 0;
- 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)
+ torch::Tensor inclusive_scan(torch::Tensor input, torch::Tensor output) {
+ const int n = input.size(0);
+ const int num_tiles = (n + TILE - 1) / TILE;
- sel = tl.arange(0, BLOCK) == BLOCK - 1
- local_agg = tl.sum(tl.where(sel, x_scan, 0.0))
+ // Grow tile state storage if needed
+ if (num_tiles > alloc_max_tiles) {
+ if (d_tile_state_storage) cudaFree(d_tile_state_storage);
+ size_t needed;
+ TileStateT::AllocationSize(num_tiles, needed);
+ d_tile_state_bytes = needed * 2;
+ cudaMalloc(&d_tile_state_storage, d_tile_state_bytes);
+ alloc_max_tiles = num_tiles;
+ }
- # ── publish PARTIAL ──────────────────────────────────────────────
- tl.store(agg_ptr + tile_id, local_agg)
- tl.atomic_xchg(status_ptr + tile_id, 1, sem='release')
+ // Initialize tile state
+ TileStateT tile_state;
+ tile_state.Init(num_tiles, d_tile_state_storage, d_tile_state_bytes);
+ init_tile_state_kernel<<<(num_tiles + 255) / 256, 256>>>(tile_state, num_tiles);
- # ── 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'
+ // Launch scan kernel
+ scan_kernel<<<num_tiles, BT>>>(
+ input.data_ptr<float>(), output.data_ptr<float>(),
+ tile_state, n);
- # ── publish COMPLETE ─────────────────────────────────────────────
- tl.store(inc_ptr + tile_id, excl + local_agg)
- tl.atomic_xchg(status_ptr + tile_id, 2, sem='release')
+ return output;
+ }
+ """
- # ── write output ─────────────────────────────────────────────────
- tl.store(out_ptr + offs, x_scan + excl, mask=mask)
+ cpp_source = r"""
+ torch::Tensor inclusive_scan(torch::Tensor input, torch::Tensor output);
+ """
- # Claim next tile
- tile_id = tl.atomic_add(counter_ptr, 1)
+ print("Compiling submission kernel...")
+ _module = load_inline(
+ name='prefix_sum_submission_v2',
+ cpp_sources=[cpp_source],
+ cuda_sources=[cuda_source],
+ functions=['inclusive_scan'],
+ verbose=False,
+ extra_cuda_cflags=['-O3', '--use_fast_math', '--threads', '4']
+ )
+ print("Done.")
- 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)
+ input_tensor, output_tensor = data
+ _module.inclusive_scan(input_tensor, output_tensor)
+ return output_tensor
scrolls · 522 diff lines total

Best evidence level for this revision: reported

JSON