submission 500141
ağaç.mp4 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 300 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-500141?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
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:79baccaa60b27c710ca32786a847e24d564217cfaeec1406472695735471d7ad
license declaredunknown
license concludedunknown
authorsağaç.mp4
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float smem_warp[NWARPS]; // warp totals for block-level scanvector-width = float4
Current (blocked float4):Kernel source
submission.py300 lines
"""
Prefix Sum - Warp-Striped DLB
@MemoryCoalesced
v17: Warp-striped IO for perfect coalescing.
Current (blocked float4):
Each warp's 32 threads read at 256-byte stride
→ 32 distinct cache lines per float4 load
→ 512 cache line fetches per warp per tile
Warp-striped:
Round v: thread lane reads float4 at warp_base + v*128 + lane*4
→ 32 threads span exactly 128 bytes = 1 cache line
→ 16 cache line fetches per warp per tile (32x reduction)
The scan is done in 16 passes, each pass covering 128 consecutive elements:
1. Load 4 elements per thread (perfectly coalesced warp load)
2. Thread-local scan of 4 elements
3. Warp exclusive prefix scan → apply carry → update running warp carry
4. Store to items[] (local inclusive prefix within warp range)
After 16 passes: items[64] holds local inclusive prefix within warp's 2048-element range.
Block scan over 8 warp totals → DLB → apply global prefix → store warp-striped (coalesced).
Expected: IO improves from ~3.4 TB/s toward ~6.5 TB/s ceiling.
Extra compute: 16 rounds × (5 SHFL + 5 FADD) = 80 extra warp ops per thread.
Net: large IO gain >> small compute overhead.
"""
from task import input_t, output_t
import torch
from torch.utils.cpp_extension import load_inline
cpp_source = """
torch::Tensor prefix_sum_cuda(torch::Tensor input, torch::Tensor output);
"""
cuda_source = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#define BLOCK_THREADS 256
#define ITEMS_PER_THREAD 64
#define TILE_SIZE (BLOCK_THREADS * ITEMS_PER_THREAD) // 16384
#define STATUS_INVALID 0
#define STATUS_PARTIAL 1
#define STATUS_INCLUSIVE 2
__device__ __forceinline__ long long pack(int status, float value) {
return ((long long)(unsigned)status << 32) | (long long)__float_as_uint(value);
}
__device__ __forceinline__ void unpack(long long packed, int &status, float &value) {
status = (int)(unsigned int)(packed >> 32);
value = __uint_as_float((unsigned int)(packed & 0xFFFFFFFFULL));
}
__device__ __forceinline__ float warp_reduce_add(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xffffffff, val, offset);
return val;
}
__device__ float warp_parallel_lookback(
volatile long long* __restrict__ descriptors, int tile_id)
{
const int lane = threadIdx.x & 31;
float prefix = 0.f;
int base = tile_id - 1;
while (true) {
int look_idx = base - lane;
int status = STATUS_INCLUSIVE;
float value = 0.f;
if (look_idx >= 0) {
long long desc;
do { desc = descriptors[look_idx]; unpack(desc, status, value); }
while (status == STATUS_INVALID);
}
unsigned imask = __ballot_sync(0xffffffff, status == STATUS_INCLUSIVE);
if (imask) {
int fi = __ffs(imask) - 1;
prefix += warp_reduce_add((lane <= fi) ? value : 0.f);
return prefix;
}
prefix += warp_reduce_add(value);
base -= 32;
if (base < 0) break;
}
return prefix;
}
// Warp-striped scan: perfectly coalesced IO, intra-warp prefix via shuffles.
// Each warp owns WARP_TILE = 2048 contiguous elements.
// Each round: 32 threads load 128 bytes (1 cache line) together.
template <bool FULL_TILE>
__global__ void __launch_bounds__(BLOCK_THREADS)
scan_kernel(
const float* __restrict__ input,
float* __restrict__ output,
volatile long long* __restrict__ descriptors,
int n)
{
constexpr int NWARPS = BLOCK_THREADS / 32; // 8
constexpr int WARP_TILE = TILE_SIZE / NWARPS; // 2048 elements per warp
constexpr int WARP_ROUNDS = ITEMS_PER_THREAD / 4; // 16 rounds of 4 elements
__shared__ float smem_warp[NWARPS]; // warp totals for block-level scan
__shared__ float s_aggregate; // tile total for DLB
__shared__ float s_prefix; // tile exclusive global prefix from DLB
const int tile_id = blockIdx.x;
const int warp_id = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int tile_start = tile_id * TILE_SIZE;
const int warp_base = tile_start + warp_id * WARP_TILE;
// in4[v * 32 + lane] = float4 at warp_base + v*128 + lane*4
// → all 32 lanes in round v load 128 consecutive bytes (1 cache line)
const float4* in4 = reinterpret_cast<const float4*>(input + warp_base);
// ----------------------------------------------------------------
// Phase 1: Warp-striped load + intra-warp inclusive prefix scan
//
// items[k] = inclusive prefix of element k within this warp's range
// ----------------------------------------------------------------
float items[ITEMS_PER_THREAD];
float warp_carry = 0.f; // running inclusive sum through completed rounds
#pragma unroll
for (int v = 0; v < WARP_ROUNDS; v++) {
float4 val;
if constexpr (FULL_TILE) {
val = __ldg(in4 + v * 32 + lane);
} else {
int gf = warp_base + v * 128 + lane * 4;
if (gf + 3 < n) {
val = __ldg(in4 + v * 32 + lane);
} else {
val.x = (gf < n) ? __ldg(input + gf ) : 0.f;
val.y = (gf+1 < n) ? __ldg(input + gf+1) : 0.f;
val.z = (gf+2 < n) ? __ldg(input + gf+2) : 0.f;
val.w = (gf+3 < n) ? __ldg(input + gf+3) : 0.f;
}
}
// Thread-local scan of 4 elements
val.y += val.x; val.z += val.y; val.w += val.z;
float thread_sum = val.w;
// Warp inclusive prefix scan on thread totals (5 shuffles)
float wp = thread_sum;
#pragma unroll
for (int off = 1; off < 32; off <<= 1) {
float t = __shfl_up_sync(0xffffffff, wp, off);
if (lane >= off) wp += t;
}
// wp = inclusive warp prefix for this lane
// exclusive warp prefix = wp - thread_sum
float carry_in = warp_carry + wp - thread_sum;
// Store local inclusive prefix: carry_in + thread-local inclusive
items[v*4+0] = val.x + carry_in;
items[v*4+1] = val.y + carry_in;
items[v*4+2] = val.z + carry_in;
items[v*4+3] = val.w + carry_in;
// Update warp_carry: broadcast lane 31's inclusive sum (= round total + prior carry)
warp_carry = __shfl_sync(0xffffffff, warp_carry + wp, 31);
}
// warp_carry = total sum of this warp's WARP_TILE elements
// ----------------------------------------------------------------
// Phase 2: Block-level exclusive prefix scan of warp totals
// Lane 0 of each warp writes to smem, lane 0 of warp 0 scans
// ----------------------------------------------------------------
if (lane == 0) smem_warp[warp_id] = warp_carry;
__syncthreads();
// Sequential scan of 8 warp totals — only lane 0 of warp 0 runs this
if (warp_id == 0 && lane == 0) {
float excl = 0.f;
#pragma unroll
for (int w = 0; w < NWARPS; w++) {
float total = smem_warp[w];
smem_warp[w] = excl; // overwrite with exclusive prefix
excl += total;
}
s_aggregate = excl; // tile total for DLB
}
__syncthreads();
float my_warp_excl = smem_warp[warp_id];
// ----------------------------------------------------------------
// Phase 3: DLB — publish partial, lookback, promote to inclusive
// ----------------------------------------------------------------
if (threadIdx.x == 0) {
long long packed = (tile_id == 0)
? pack(STATUS_INCLUSIVE, s_aggregate)
: pack(STATUS_PARTIAL, s_aggregate);
descriptors[tile_id] = packed;
__threadfence();
}
__syncthreads();
if (tile_id > 0 && warp_id == 0) {
float lbprefix = warp_parallel_lookback(descriptors, tile_id);
if (lane == 0) {
s_prefix = lbprefix;
descriptors[tile_id] = pack(STATUS_INCLUSIVE, lbprefix + s_aggregate);
__threadfence();
}
} else if (tile_id == 0 && threadIdx.x == 0) {
s_prefix = 0.f;
}
__syncthreads();
// ----------------------------------------------------------------
// Phase 4: Apply global prefix + warp prefix, store warp-striped
// (same pattern as load — perfectly coalesced)
// ----------------------------------------------------------------
float total_add = s_prefix + my_warp_excl;
float4* out4 = reinterpret_cast<float4*>(output + warp_base);
#pragma unroll
for (int v = 0; v < WARP_ROUNDS; v++) {
float4 val = {
items[v*4+0] + total_add,
items[v*4+1] + total_add,
items[v*4+2] + total_add,
items[v*4+3] + total_add,
};
if constexpr (FULL_TILE) {
out4[v * 32 + lane] = val;
} else {
int gf = warp_base + v * 128 + lane * 4;
if (gf + 3 < n) {
out4[v * 32 + lane] = val;
} else {
if (gf < n) output[gf ] = val.x;
if (gf+1 < n) output[gf+1] = val.y;
if (gf+2 < n) output[gf+2] = val.z;
if (gf+3 < n) output[gf+3] = val.w;
}
}
}
}
template __global__ void scan_kernel<true>(const float*, float*, volatile long long*, int);
template __global__ void scan_kernel<false>(const float*, float*, volatile long long*, int);
torch::Tensor prefix_sum_cuda(torch::Tensor input, torch::Tensor output) {
const int n = input.numel();
const int num_tiles = (n + TILE_SIZE - 1) / TILE_SIZE;
auto desc = torch::zeros({num_tiles},
torch::TensorOptions().dtype(torch::kInt64).device(input.device()));
auto* d_in = input.data_ptr<float>();
auto* d_out = output.data_ptr<float>();
auto* d_desc = (volatile long long*)desc.data_ptr<int64_t>();
if (n % TILE_SIZE == 0) {
scan_kernel<true><<<num_tiles, BLOCK_THREADS>>>(d_in, d_out, d_desc, n);
} else {
scan_kernel<false><<<num_tiles, BLOCK_THREADS>>>(d_in, d_out, d_desc, n);
}
return output;
}
"""
module = None
def get_module():
global module
if module is None:
module = load_inline(
name='warp_striped_v17',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['prefix_sum_cuda'],
verbose=False,
extra_cuda_cflags=[
'-O3',
'--use_fast_math',
'-std=c++17',
'--ptxas-options=-v',
],
)
return module
def custom_kernel(data: input_t) -> output_t:
inp, out = data
mod = get_module()
mod.prefix_sum_cuda(inp, out)
return out
scrolls · 300 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 500137.
"""- Prefix Sum - DLB with Warp-Parallel Lookback+ Prefix Sum - Warp-Striped DLB@MemoryCoalesced- v16: Eliminate boundary-check ISETPs from hot path.+ v17: Warp-striped IO for perfect coalescing.- PTX analysis showed 147 ISETP instructions — almost all from:- if (gf + 3 < n) ... else scalar fallback- on every load and store. For n=268,435,456 = 16384 * TILE_SIZE,- every tile is full so these checks are dead code in the contest.+ Current (blocked float4):+ Each warp's 32 threads read at 256-byte stride+ → 32 distinct cache lines per float4 load+ → 512 cache line fetches per warp per tile- Fix: template on bool FULL_TILE.- - FULL_TILE=true: pure float4 loads/stores, zero ISETPs, zero branches- - FULL_TILE=false: original scalar fallback for tail tile- Dispatch: if n % TILE_SIZE == 0, all tiles take FULL_TILE=true path.+ Warp-striped:+ Round v: thread lane reads float4 at warp_base + v*128 + lane*4+ → 32 threads span exactly 128 bytes = 1 cache line+ → 16 cache line fetches per warp per tile (32x reduction)++ The scan is done in 16 passes, each pass covering 128 consecutive elements:+ 1. Load 4 elements per thread (perfectly coalesced warp load)+ 2. Thread-local scan of 4 elements+ 3. Warp exclusive prefix scan → apply carry → update running warp carry+ 4. Store to items[] (local inclusive prefix within warp range)++ After 16 passes: items[64] holds local inclusive prefix within warp's 2048-element range.+ Block scan over 8 warp totals → DLB → apply global prefix → store warp-striped (coalesced).++ Expected: IO improves from ~3.4 TB/s toward ~6.5 TB/s ceiling.+ Extra compute: 16 rounds × (5 SHFL + 5 FADD) = 80 extra warp ops per thread.+ Net: large IO gain >> small compute overhead."""from task import input_t, output_t⋯ 24 unchanged linesvalue = __uint_as_float((unsigned int)(packed & 0xFFFFFFFFULL));}- __device__ __forceinline__ float warp_inclusive_scan(float val) {- #pragma unroll- for (int offset = 1; offset < 32; offset <<= 1) {- float tmp = __shfl_up_sync(0xffffffff, val, offset);- if ((threadIdx.x & 31) >= offset) val += tmp;- }- return val;- }-- __device__ float block_inclusive_scan(float val, float* smem_scan) {- const int lane = threadIdx.x & 31;- const int warp = threadIdx.x >> 5;- const int nwarps = BLOCK_THREADS / 32;-- val = warp_inclusive_scan(val);- if (lane == 31) smem_scan[warp] = val;- __syncthreads();-- if (warp == 0) {- float w = (lane < nwarps) ? smem_scan[lane] : 0.0f;- w = warp_inclusive_scan(w);- if (lane < nwarps) smem_scan[lane] = w;- }- __syncthreads();-- if (warp > 0) val += smem_scan[warp - 1];- return val;- }-__device__ __forceinline__ float warp_reduce_add(float val) {#pragma unrollfor (int offset = 16; offset > 0; offset >>= 1)⋯ 2 unchanged lines}__device__ float warp_parallel_lookback(- volatile long long* __restrict__ descriptors,- int tile_id)+ volatile long long* __restrict__ descriptors, int tile_id){const int lane = threadIdx.x & 31;- float prefix = 0.0f;+ float prefix = 0.f;int base = tile_id - 1;-while (true) {int look_idx = base - lane;int status = STATUS_INCLUSIVE;- float value = 0.0f;-+ float value = 0.f;if (look_idx >= 0) {long long desc;- do {- desc = descriptors[look_idx];- unpack(desc, status, value);- } while (status == STATUS_INVALID);+ do { desc = descriptors[look_idx]; unpack(desc, status, value); }+ while (status == STATUS_INVALID);}-- unsigned inclusive_mask = __ballot_sync(0xffffffff, status == STATUS_INCLUSIVE);- if (inclusive_mask) {- int first_inc = __ffs(inclusive_mask) - 1;- float contrib = (lane <= first_inc) ? value : 0.0f;- prefix += warp_reduce_add(contrib);+ unsigned imask = __ballot_sync(0xffffffff, status == STATUS_INCLUSIVE);+ if (imask) {+ int fi = __ffs(imask) - 1;+ prefix += warp_reduce_add((lane <= fi) ? value : 0.f);return prefix;}-prefix += warp_reduce_add(value);base -= 32;if (base < 0) break;⋯ 1 unchanged linesreturn prefix;}+ // Warp-striped scan: perfectly coalesced IO, intra-warp prefix via shuffles.+ // Each warp owns WARP_TILE = 2048 contiguous elements.+ // Each round: 32 threads load 128 bytes (1 cache line) together.template <bool FULL_TILE>__global__ void __launch_bounds__(BLOCK_THREADS)scan_kernel(⋯ 2 unchanged linesvolatile long long* __restrict__ descriptors,int n){- __shared__ float smem_scan[8];- __shared__ float s_aggregate;- __shared__ float s_prefix;+ constexpr int NWARPS = BLOCK_THREADS / 32; // 8+ constexpr int WARP_TILE = TILE_SIZE / NWARPS; // 2048 elements per warp+ constexpr int WARP_ROUNDS = ITEMS_PER_THREAD / 4; // 16 rounds of 4 elements- const int tile_id = blockIdx.x;- const int tile_start = tile_id * TILE_SIZE;- const int thread_start = tile_start + threadIdx.x * ITEMS_PER_THREAD;- constexpr int VEC = ITEMS_PER_THREAD / 4;+ __shared__ float smem_warp[NWARPS]; // warp totals for block-level scan+ __shared__ float s_aggregate; // tile total for DLB+ __shared__ float s_prefix; // tile exclusive global prefix from DLB+ const int tile_id = blockIdx.x;+ const int warp_id = threadIdx.x >> 5;+ const int lane = threadIdx.x & 31;+ const int tile_start = tile_id * TILE_SIZE;+ const int warp_base = tile_start + warp_id * WARP_TILE;++ // in4[v * 32 + lane] = float4 at warp_base + v*128 + lane*4+ // → all 32 lanes in round v load 128 consecutive bytes (1 cache line)+ const float4* in4 = reinterpret_cast<const float4*>(input + warp_base);+// ----------------------------------------------------------------- // 1. Load — float4 throughout, boundary checks only in tail tile+ // Phase 1: Warp-striped load + intra-warp inclusive prefix scan+ //+ // items[k] = inclusive prefix of element k within this warp's range// ----------------------------------------------------------------float items[ITEMS_PER_THREAD];- const float4* in4 = reinterpret_cast<const float4*>(input + thread_start);+ float warp_carry = 0.f; // running inclusive sum through completed rounds#pragma unroll- for (int v = 0; v < VEC; v++) {+ for (int v = 0; v < WARP_ROUNDS; v++) {float4 val;if constexpr (FULL_TILE) {- // No branch, no ISETP — compiler emits pure LDG.E.128- val = __ldg(in4 + v);+ val = __ldg(in4 + v * 32 + lane);} else {- int gf = thread_start + v * 4;+ int gf = warp_base + v * 128 + lane * 4;if (gf + 3 < n) {- val = __ldg(in4 + v);+ val = __ldg(in4 + v * 32 + lane);} else {- val.x = (gf < n) ? __ldg(input + gf ) : 0.0f;- val.y = (gf + 1 < n) ? __ldg(input + gf + 1) : 0.0f;- val.z = (gf + 2 < n) ? __ldg(input + gf + 2) : 0.0f;- val.w = (gf + 3 < n) ? __ldg(input + gf + 3) : 0.0f;+ val.x = (gf < n) ? __ldg(input + gf ) : 0.f;+ val.y = (gf+1 < n) ? __ldg(input + gf+1) : 0.f;+ val.z = (gf+2 < n) ? __ldg(input + gf+2) : 0.f;+ val.w = (gf+3 < n) ? __ldg(input + gf+3) : 0.f;}}- items[v*4+0] = val.x; items[v*4+1] = val.y;- items[v*4+2] = val.z; items[v*4+3] = val.w;++ // Thread-local scan of 4 elements+ val.y += val.x; val.z += val.y; val.w += val.z;+ float thread_sum = val.w;++ // Warp inclusive prefix scan on thread totals (5 shuffles)+ float wp = thread_sum;+ #pragma unroll+ for (int off = 1; off < 32; off <<= 1) {+ float t = __shfl_up_sync(0xffffffff, wp, off);+ if (lane >= off) wp += t;+ }+ // wp = inclusive warp prefix for this lane+ // exclusive warp prefix = wp - thread_sum+ float carry_in = warp_carry + wp - thread_sum;++ // Store local inclusive prefix: carry_in + thread-local inclusive+ items[v*4+0] = val.x + carry_in;+ items[v*4+1] = val.y + carry_in;+ items[v*4+2] = val.z + carry_in;+ items[v*4+3] = val.w + carry_in;++ // Update warp_carry: broadcast lane 31's inclusive sum (= round total + prior carry)+ warp_carry = __shfl_sync(0xffffffff, warp_carry + wp, 31);}+ // warp_carry = total sum of this warp's WARP_TILE elements// ----------------------------------------------------------------- // 2. Thread-local scan+ // Phase 2: Block-level exclusive prefix scan of warp totals+ // Lane 0 of each warp writes to smem, lane 0 of warp 0 scans// ----------------------------------------------------------------- #pragma unroll- for (int i = 1; i < ITEMS_PER_THREAD; i++) items[i] += items[i-1];- float thread_sum = items[ITEMS_PER_THREAD - 1];+ if (lane == 0) smem_warp[warp_id] = warp_carry;+ __syncthreads();- // ----------------------------------------------------------------- // 3. Block scan- // ----------------------------------------------------------------- float block_result = block_inclusive_scan(thread_sum, smem_scan);+ // Sequential scan of 8 warp totals — only lane 0 of warp 0 runs this+ if (warp_id == 0 && lane == 0) {+ float excl = 0.f;+ #pragma unroll+ for (int w = 0; w < NWARPS; w++) {+ float total = smem_warp[w];+ smem_warp[w] = excl; // overwrite with exclusive prefix+ excl += total;+ }+ s_aggregate = excl; // tile total for DLB+ }+ __syncthreads();+ float my_warp_excl = smem_warp[warp_id];+// ----------------------------------------------------------------- // 4. Publish descriptor+ // Phase 3: DLB — publish partial, lookback, promote to inclusive// ----------------------------------------------------------------- if (threadIdx.x == BLOCK_THREADS - 1) {- s_aggregate = block_result;+ if (threadIdx.x == 0) {long long packed = (tile_id == 0)- ? pack(STATUS_INCLUSIVE, block_result)- : pack(STATUS_PARTIAL, block_result);+ ? pack(STATUS_INCLUSIVE, s_aggregate)+ : pack(STATUS_PARTIAL, s_aggregate);descriptors[tile_id] = packed;__threadfence();}__syncthreads();- // ----------------------------------------------------------------- // 5. Warp 0: lookback + promote- // ----------------------------------------------------------------- if (tile_id > 0 && (threadIdx.x >> 5) == 0) {+ if (tile_id > 0 && warp_id == 0) {float lbprefix = warp_parallel_lookback(descriptors, tile_id);- if (threadIdx.x == 0) {+ if (lane == 0) {s_prefix = lbprefix;descriptors[tile_id] = pack(STATUS_INCLUSIVE, lbprefix + s_aggregate);__threadfence();}} else if (tile_id == 0 && threadIdx.x == 0) {- s_prefix = 0.0f;+ s_prefix = 0.f;}__syncthreads();// ----------------------------------------------------------------- // 6. Apply prefix + store+ // Phase 4: Apply global prefix + warp prefix, store warp-striped+ // (same pattern as load — perfectly coalesced)// ----------------------------------------------------------------- float my_prefix = s_prefix + (block_result - thread_sum);- #pragma unroll- for (int i = 0; i < ITEMS_PER_THREAD; i++) items[i] += my_prefix;+ float total_add = s_prefix + my_warp_excl;+ float4* out4 = reinterpret_cast<float4*>(output + warp_base);- float4* out4 = reinterpret_cast<float4*>(output + thread_start);#pragma unroll- for (int v = 0; v < VEC; v++) {- float4 val = { items[v*4+0], items[v*4+1], items[v*4+2], items[v*4+3] };+ for (int v = 0; v < WARP_ROUNDS; v++) {+ float4 val = {+ items[v*4+0] + total_add,+ items[v*4+1] + total_add,+ items[v*4+2] + total_add,+ items[v*4+3] + total_add,+ };if constexpr (FULL_TILE) {- out4[v] = val;+ out4[v * 32 + lane] = val;} else {- int gf = thread_start + v * 4;+ int gf = warp_base + v * 128 + lane * 4;if (gf + 3 < n) {- out4[v] = val;+ out4[v * 32 + lane] = val;} else {- if (gf < n) output[gf ] = val.x;- if (gf + 1 < n) output[gf + 1] = val.y;- if (gf + 2 < n) output[gf + 2] = val.z;- if (gf + 3 < n) output[gf + 3] = val.w;+ if (gf < n) output[gf ] = val.x;+ if (gf+1 < n) output[gf+1] = val.y;+ if (gf+2 < n) output[gf+2] = val.z;+ if (gf+3 < n) output[gf+3] = val.w;}}}}- // Explicit instantiations so both are compiledtemplate __global__ void scan_kernel<true>(const float*, float*, volatile long long*, int);template __global__ void scan_kernel<false>(const float*, float*, volatile long long*, int);⋯ 4 unchanged linesauto desc = torch::zeros({num_tiles},torch::TensorOptions().dtype(torch::kInt64).device(input.device()));- auto* d_in = input.data_ptr<float>();- auto* d_out = output.data_ptr<float>();+ auto* d_in = input.data_ptr<float>();+ auto* d_out = output.data_ptr<float>();auto* d_desc = (volatile long long*)desc.data_ptr<int64_t>();if (n % TILE_SIZE == 0) {- // All tiles full — zero boundary checksscan_kernel<true><<<num_tiles, BLOCK_THREADS>>>(d_in, d_out, d_desc, n);} else {- // Last tile is partial — use safe path for all tiles- // (could split: full path for [0, num_tiles-1], safe for last — but- // this case doesn't occur in the contest so not worth the complexity)scan_kernel<false><<<num_tiles, BLOCK_THREADS>>>(d_in, d_out, d_desc, n);}⋯ 7 unchanged linesglobal moduleif module is None:module = load_inline(- name='dlb_scan_v16',+ name='warp_striped_v17',cpp_sources=cpp_source,cuda_sources=cuda_source,functions=['prefix_sum_cuda'],⋯ 2 unchanged lines'-O3','--use_fast_math','-std=c++17',+ '--ptxas-options=-v',],)return module
scrolls · 375 diff lines total
Best evidence level for this revision: reported
JSON