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
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_inlinefrom 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