Skip to content
KernelIndex
Search⌘K

submission 522494

tomaszki · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

prefixsum_dualwarp_tma.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-522494?include=source"
interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
Inclusive prefix sumsuite of 11 cases
NVIDIA H100
761.7µs
#2 of 23
2026-03-09

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b593122bfce1e9310f29c5c7eff07b70892bd6c3f014033ce004dd6e4ed053e1
license declaredunknown
license concludedunknown
authorstomaszki
imported2026-08-15

Techniques

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

async-copyptx::cp_async_bulk(
mbarrierptx::mbarrier_init(&load_mbar[s], 1);
shared-memoryextern __shared__ float4 stage_data[];
stages = 4GLOBAL_STAGES = 4
vector-width = float4__device__ __forceinline__ float4 ld_relaxed_gpu_v4_f32(const float* src) {

Kernel source

prefixsum_dualwarp_tma.py438 lines
#!POPCORN leaderboard prefixsum_v2
#!POPCORN gpu H100

import torch
from torch.utils.cpp_extension import load_inline
from typing import List
from task import input_t, output_t

GLOBAL_STAGES = 4
LOCAL_STAGES = 111
FIXED_N = 268_435_456

if GLOBAL_STAGES < 1:
    raise ValueError("GLOBAL_STAGES must be >= 1")
if LOCAL_STAGES < 1:
    raise ValueError("LOCAL_STAGES must be >= 1")
if FIXED_N < 1:
    raise ValueError("FIXED_N must be >= 1")

# Per block per round: LOCAL_STAGES * 32 lanes * 4 floats.
VALUES_PER_BLOCK_ROUND = LOCAL_STAGES * 128
ROUND_SPAN = 128 * VALUES_PER_BLOCK_ROUND


prefixsum_cuda_source = r"""
#include <cuda_runtime.h>
#include <cuda/ptx>

#ifndef GLOBAL_STAGES
#define GLOBAL_STAGES 2
#endif

#ifndef LOCAL_STAGES
#define LOCAL_STAGES 8
#endif

#ifndef FIXED_N
#define FIXED_N 268435456
#endif

static_assert(GLOBAL_STAGES >= 1, "GLOBAL_STAGES must be >= 1");
static_assert(LOCAL_STAGES >= 1, "LOCAL_STAGES must be >= 1");
static_assert(FIXED_N >= 1, "FIXED_N must be >= 1");
static_assert((FIXED_N % 128) == 0, "FIXED_N must be divisible by 128");

namespace ptx = cuda::ptx;

__device__ __forceinline__ float ld_relaxed_gpu_f32(const float* src) {
    float out;
    asm volatile(
        "ld.relaxed.gpu.f32 %0, [%1];"
        : "=f"(out)
        : "l"(src)
        : "memory"
    );
    return out;
}

__device__ __forceinline__ void st_relaxed_gpu_f32(float* dst, float value) {
    asm volatile(
        "st.relaxed.gpu.f32 [%0], %1;"
        :
        : "l"(dst), "f"(value)
        : "memory"
    );
}

__device__ __forceinline__ float4 ld_relaxed_gpu_v4_f32(const float* src) {
    float4 out;
    asm volatile(
        "ld.relaxed.gpu.v4.f32 {%0, %1, %2, %3}, [%4];"
        : "=f"(out.x), "=f"(out.y), "=f"(out.z), "=f"(out.w)
        : "l"(src)
        : "memory"
    );
    return out;
}

__device__ __forceinline__ void st_relaxed_gpu_v4_f32(float* dst, const float4& value) {
    asm volatile(
        "st.relaxed.gpu.v4.f32 [%0], {%1, %2, %3, %4};"
        :
        : "l"(dst), "f"(value.x), "f"(value.y), "f"(value.z), "f"(value.w)
        : "memory"
    );
}

__device__ __forceinline__ unsigned int with_flag(unsigned int bits) {
    return bits | 1u;
}

__device__ __forceinline__ unsigned int clear_flag(unsigned int bits) {
    return bits & ~1u;
}

__device__ __forceinline__ bool flag_is_set(unsigned int bits) {
    return (bits & 1u) != 0u;
}

__device__ __forceinline__ bool is_elected_in_warp(const int target_warp) {
    const unsigned int warp_id = static_cast<unsigned int>(threadIdx.x >> 5);
    const unsigned int uniform_warp_id = __shfl_sync(0xffffffffu, warp_id, 0);
    return (uniform_warp_id == static_cast<unsigned int>(target_warp)) && ptx::elect_sync(0xffffffffu);
}

__global__ void prefixsum_kernel(const float* input, float* output, float* block_sums) {
    const int block_idx = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;

    const int worker_blocks = 128;
    const int coordinator_block = 128;
    const int starter_warp = 0;
    const int finalizer_warp = 1;
    const int warp_size = 32;
    const int values_per_thread = 4;
    const int values_per_local = warp_size * values_per_thread;  // 128 floats
    const int values_per_block_round = LOCAL_STAGES * values_per_local;
    const int slot_float4s = LOCAL_STAGES * warp_size;
    const int round_span = worker_blocks * values_per_block_round;
    const int rounds = (FIXED_N + round_span - 1) / round_span;
    const unsigned int full_mask = 0xffffffffu;

    if (block_idx < worker_blocks) {
        extern __shared__ float4 stage_data[];
        __shared__ uint64_t load_mbar[GLOBAL_STAGES];
        __shared__ uint64_t done_mbar[GLOBAL_STAGES];

        if (tid == 0) {
            #pragma unroll
            for (int s = 0; s < GLOBAL_STAGES; ++s) {
                ptx::mbarrier_init(&load_mbar[s], 1);
                ptx::mbarrier_init(&done_mbar[s], 1);
            }
        }
        __syncthreads();

        if (warp == finalizer_warp) {
            const int warmup = (rounds < GLOBAL_STAGES) ? rounds : GLOBAL_STAGES;
            #pragma unroll
            for (int r = 0; r < warmup; ++r) {
                const int slot = r;
                const int src_base_idx = r * round_span + block_idx * values_per_block_round;
                int valid_values = FIXED_N - src_base_idx;
                if (valid_values < 0) valid_values = 0;
                if (valid_values > values_per_block_round) valid_values = values_per_block_round;
                const int valid_local_stages = valid_values / values_per_local;
                const uint32_t valid_bytes = static_cast<uint32_t>(valid_local_stages * warp_size * sizeof(float4));
                if (valid_local_stages > 0) {
                    if (is_elected_in_warp(finalizer_warp)) {
                        float4* dst = &stage_data[slot * slot_float4s];
                        const float4* src = reinterpret_cast<const float4*>(input + src_base_idx);
                        ptx::cp_async_bulk(
                            ptx::space_shared, ptx::space_global,
                            dst, src, valid_bytes, &load_mbar[slot]
                        );
                        (void)ptx::mbarrier_arrive_expect_tx(
                            ptx::sem_release, ptx::scope_cta, ptx::space_shared,
                            &load_mbar[slot], valid_bytes
                        );
                    }
                } else {
                    if (is_elected_in_warp(finalizer_warp)) {
                        (void)ptx::mbarrier_arrive(
                            ptx::sem_release, ptx::scope_cta, ptx::space_shared, &load_mbar[slot]
                        );
                    }
                }
            }

            for (int round = 0; round < rounds; ++round) {
                const int slot = round % GLOBAL_STAGES;
                const uint32_t phase = static_cast<uint32_t>((round / GLOBAL_STAGES) & 1);
                while (!ptx::mbarrier_try_wait_parity(
                    ptx::sem_acquire, ptx::scope_cta, &done_mbar[slot], phase
                )) {}

                const float* mailbox_ptr = &block_sums[slot * worker_blocks + block_idx];
                unsigned int bits = 0u;
                do {
                    bits = __float_as_uint(ld_relaxed_gpu_f32(mailbox_ptr));
                } while (flag_is_set(bits));
                const float block_base = __uint_as_float(clear_flag(bits));

                const int out_base_idx = round * round_span + block_idx * values_per_block_round;
                int valid_values = FIXED_N - out_base_idx;
                if (valid_values < 0) valid_values = 0;
                if (valid_values > values_per_block_round) valid_values = values_per_block_round;
                const int valid_local_stages = valid_values / values_per_local;
                
                for (int ls = 0; ls < valid_local_stages; ++ls) {
                    const int smem_idx = slot * slot_float4s + ls * warp_size + lane;
                    float4 out4 = stage_data[smem_idx];
                    out4.x += block_base;
                    out4.y += block_base;
                    out4.z += block_base;
                    out4.w += block_base;

                    const int dst_idx = out_base_idx + ls * values_per_local + lane * values_per_thread;
                    reinterpret_cast<float4*>(output)[dst_idx / 4] = out4;
                }

                const int fetch_round = round + GLOBAL_STAGES;
                if (fetch_round < rounds) {
                    const int src_base_idx = fetch_round * round_span + block_idx * values_per_block_round;
                    int next_valid_values = FIXED_N - src_base_idx;
                    if (next_valid_values < 0) next_valid_values = 0;
                    if (next_valid_values > values_per_block_round) next_valid_values = values_per_block_round;
                    const int next_valid_local_stages = next_valid_values / values_per_local;
                    const uint32_t next_valid_bytes =
                        static_cast<uint32_t>(next_valid_local_stages * warp_size * sizeof(float4));
                    if (next_valid_local_stages > 0) {
                        if (is_elected_in_warp(finalizer_warp)) {
                            float4* dst = &stage_data[slot * slot_float4s];
                            const float4* src = reinterpret_cast<const float4*>(input + src_base_idx);
                            ptx::cp_async_bulk(
                                ptx::space_shared, ptx::space_global,
                                dst, src, next_valid_bytes, &load_mbar[slot]
                            );
                            (void)ptx::mbarrier_arrive_expect_tx(
                                ptx::sem_release, ptx::scope_cta, ptx::space_shared,
                                &load_mbar[slot], next_valid_bytes
                            );
                        }
                    } else {
                        if (is_elected_in_warp(finalizer_warp)) {
                            (void)ptx::mbarrier_arrive(
                                ptx::sem_release, ptx::scope_cta, ptx::space_shared, &load_mbar[slot]
                            );
                        }
                    }
                }
            }
        } else if (warp == starter_warp) {
            for (int round = 0; round < rounds; ++round) {
                const int slot = round % GLOBAL_STAGES;
                const uint32_t phase = static_cast<uint32_t>((round / GLOBAL_STAGES) & 1);
                while (!ptx::mbarrier_try_wait_parity(
                    ptx::sem_acquire, ptx::scope_cta, &load_mbar[slot], phase
                )) {}

                const int round_base_idx = round * round_span + block_idx * values_per_block_round;
                int valid_values = FIXED_N - round_base_idx;
                if (valid_values < 0) valid_values = 0;
                if (valid_values > values_per_block_round) valid_values = values_per_block_round;
                const int valid_local_stages = valid_values / values_per_local;
                const int valid_float4 = valid_local_stages * warp_size;
                int thread_items = valid_float4 - lane * LOCAL_STAGES;
                if (thread_items < 0) thread_items = 0;
                if (thread_items > LOCAL_STAGES) thread_items = LOCAL_STAGES;

                const int slot_base = slot * slot_float4s;
                const int thread_base = slot_base + lane * LOCAL_STAGES;

                // Phase 1: each thread scans its contiguous LOCAL_STAGES chunk and stores partials.
                float thread_running = 0.0f;
                for (int i = 0; i < thread_items; ++i) {
                    const int smem_idx = thread_base + i;
                    const float4 in4 = stage_data[smem_idx];

                    const float p0 = in4.x;
                    const float p1 = p0 + in4.y;
                    const float p2 = p1 + in4.z;
                    const float p3 = p2 + in4.w;

                    float4 partial4;
                    partial4.x = p0 + thread_running;
                    partial4.y = p1 + thread_running;
                    partial4.z = p2 + thread_running;
                    partial4.w = p3 + thread_running;
                    stage_data[smem_idx] = partial4;
                    thread_running += p3;
                }

                // Phase 2: warp scan over per-thread totals to get thread-level exclusive carry.
                float thread_inclusive = thread_running;
                #pragma unroll
                for (int offset = 1; offset < warp_size; offset <<= 1) {
                    const float y = __shfl_up_sync(full_mask, thread_inclusive, offset);
                    if (lane >= offset) {
                        thread_inclusive += y;
                    }
                }
                const float thread_exclusive = thread_inclusive - thread_running;

                // Publish mailbox immediately after phase 2 once block total is known.
                if (lane == warp_size - 1) {
                    const unsigned int tagged_bits = with_flag(__float_as_uint(thread_inclusive));
                    st_relaxed_gpu_f32(
                        &block_sums[slot * worker_blocks + block_idx],
                        __uint_as_float(tagged_bits)
                    );
                }

                // Phase 3: reload contiguous chunk and add thread carry.
                for (int i = 0; i < thread_items; ++i) {
                    const int smem_idx = thread_base + i;
                    float4 out4 = stage_data[smem_idx];
                    out4.x += thread_exclusive;
                    out4.y += thread_exclusive;
                    out4.z += thread_exclusive;
                    out4.w += thread_exclusive;
                    stage_data[smem_idx] = out4;
                }

                if (lane == warp_size - 1) {
                    (void)ptx::mbarrier_arrive(
                        ptx::sem_release, ptx::scope_cta, ptx::space_shared, &done_mbar[slot]
                    );
                }
            }
        }
        return;
    }

    if (block_idx == coordinator_block) {
        if (tid >= 32) {
            return;
        }

        float running_sum = 0.0f;
        for (int round = 0; round < rounds; ++round) {
            const int slot = round % GLOBAL_STAGES;
            const float* slot_base_ptr = &block_sums[slot * worker_blocks];

            float4 tagged4;
            unsigned int bits0 = 0u;
            unsigned int bits1 = 0u;
            unsigned int bits2 = 0u;
            unsigned int bits3 = 0u;
            bool all_ready = false;
            do {
                const float* slot_ptr = &slot_base_ptr[4 * lane];
                tagged4 = ld_relaxed_gpu_v4_f32(slot_ptr);

                bits0 = __float_as_uint(tagged4.x);
                bits1 = __float_as_uint(tagged4.y);
                bits2 = __float_as_uint(tagged4.z);
                bits3 = __float_as_uint(tagged4.w);

                const bool ready4 =
                    flag_is_set(bits0) && flag_is_set(bits1) &&
                    flag_is_set(bits2) && flag_is_set(bits3);
                all_ready = __all_sync(full_mask, ready4);
            } while (!all_ready);

            const float s0 = __uint_as_float(clear_flag(bits0));
            const float s1 = __uint_as_float(clear_flag(bits1));
            const float s2 = __uint_as_float(clear_flag(bits2));
            const float s3 = __uint_as_float(clear_flag(bits3));

            const float lane_total = s0 + s1 + s2 + s3;
            float warp_scan = lane_total;

            #pragma unroll
            for (int offset = 1; offset < 32; offset <<= 1) {
                const float y = __shfl_up_sync(full_mask, warp_scan, offset);
                if (lane >= offset) {
                    warp_scan += y;
                }
            }

            const float lane_base = warp_scan - lane_total;
            const float o0 = running_sum + lane_base;
            const float o1 = o0 + s0;
            const float o2 = o1 + s1;
            const float o3 = o2 + s2;

            float4 out4;
            out4.x = __uint_as_float(clear_flag(__float_as_uint(o0)));
            out4.y = __uint_as_float(clear_flag(__float_as_uint(o1)));
            out4.z = __uint_as_float(clear_flag(__float_as_uint(o2)));
            out4.w = __uint_as_float(clear_flag(__float_as_uint(o3)));
            st_relaxed_gpu_v4_f32(&block_sums[slot * worker_blocks + 4 * lane], out4);

            const float round_total = __shfl_sync(full_mask, warp_scan, 31);
            running_sum += round_total;
        }
    }
}

torch::Tensor prefixsum_cuda(torch::Tensor input) {
    const int size = input.numel();
    auto output = torch::empty({size}, input.options());
    auto block_sums = torch::zeros({128 * GLOBAL_STAGES}, input.options());
    constexpr int threads_per_block = 64;
    const size_t shared_bytes =
        static_cast<size_t>(GLOBAL_STAGES) * LOCAL_STAGES * 32 * sizeof(float4);
    cudaFuncSetAttribute(
        prefixsum_kernel,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        static_cast<int>(shared_bytes)
    );
    cudaFuncSetAttribute(
        prefixsum_kernel,
        cudaFuncAttributePreferredSharedMemoryCarveout,
        100
    );
    prefixsum_kernel<<<129, threads_per_block, shared_bytes>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        block_sums.data_ptr<float>()
    );
    return output;
}
"""


prefixsum_cpp_source = """
#include <torch/extension.h>

torch::Tensor prefixsum_cuda(torch::Tensor a);
"""

prefixsum_module = load_inline(
    name="prefixsum_cuda_dualwarp_tma",
    cpp_sources=prefixsum_cpp_source,
    cuda_sources=prefixsum_cuda_source,
    functions=["prefixsum_cuda"],
    verbose=True,
    extra_cuda_cflags=[
        "-O3",
        f"-DGLOBAL_STAGES={GLOBAL_STAGES}",
        f"-DLOCAL_STAGES={LOCAL_STAGES}",
        f"-DFIXED_N={FIXED_N}",
        "-gencode=arch=compute_90,code=sm_90",
    ],
    extra_ldflags=["-lcuda"],
)


def custom_kernel(data: input_t) -> output_t:
    data = data[0]
    if data.numel() != FIXED_N:
        return torch.cumsum(data, dim=0)
    return prefixsum_module.prefixsum_cuda(data)
scrolls · 438 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 183714.

+ #!POPCORN leaderboard prefixsum_v2
+ #!POPCORN gpu H100
+
import torch
+ from torch.utils.cpp_extension import load_inline
+ from typing import List
from task import input_t, output_t
+ GLOBAL_STAGES = 4
+ LOCAL_STAGES = 111
+ FIXED_N = 268_435_456
+ if GLOBAL_STAGES < 1:
+ raise ValueError("GLOBAL_STAGES must be >= 1")
+ if LOCAL_STAGES < 1:
+ raise ValueError("LOCAL_STAGES must be >= 1")
+ if FIXED_N < 1:
+ raise ValueError("FIXED_N must be >= 1")
+
+ # Per block per round: LOCAL_STAGES * 32 lanes * 4 floats.
+ VALUES_PER_BLOCK_ROUND = LOCAL_STAGES * 128
+ ROUND_SPAN = 128 * VALUES_PER_BLOCK_ROUND
+
+
+ prefixsum_cuda_source = r"""
+ #include <cuda_runtime.h>
+ #include <cuda/ptx>
+
+ #ifndef GLOBAL_STAGES
+ #define GLOBAL_STAGES 2
+ #endif
+
+ #ifndef LOCAL_STAGES
+ #define LOCAL_STAGES 8
+ #endif
+
+ #ifndef FIXED_N
+ #define FIXED_N 268435456
+ #endif
+
+ static_assert(GLOBAL_STAGES >= 1, "GLOBAL_STAGES must be >= 1");
+ static_assert(LOCAL_STAGES >= 1, "LOCAL_STAGES must be >= 1");
+ static_assert(FIXED_N >= 1, "FIXED_N must be >= 1");
+ static_assert((FIXED_N % 128) == 0, "FIXED_N must be divisible by 128");
+
+ namespace ptx = cuda::ptx;
+
+ __device__ __forceinline__ float ld_relaxed_gpu_f32(const float* src) {
+ float out;
+ asm volatile(
+ "ld.relaxed.gpu.f32 %0, [%1];"
+ : "=f"(out)
+ : "l"(src)
+ : "memory"
+ );
+ return out;
+ }
+
+ __device__ __forceinline__ void st_relaxed_gpu_f32(float* dst, float value) {
+ asm volatile(
+ "st.relaxed.gpu.f32 [%0], %1;"
+ :
+ : "l"(dst), "f"(value)
+ : "memory"
+ );
+ }
+
+ __device__ __forceinline__ float4 ld_relaxed_gpu_v4_f32(const float* src) {
+ float4 out;
+ asm volatile(
+ "ld.relaxed.gpu.v4.f32 {%0, %1, %2, %3}, [%4];"
+ : "=f"(out.x), "=f"(out.y), "=f"(out.z), "=f"(out.w)
+ : "l"(src)
+ : "memory"
+ );
+ return out;
+ }
+
+ __device__ __forceinline__ void st_relaxed_gpu_v4_f32(float* dst, const float4& value) {
+ asm volatile(
+ "st.relaxed.gpu.v4.f32 [%0], {%1, %2, %3, %4};"
+ :
+ : "l"(dst), "f"(value.x), "f"(value.y), "f"(value.z), "f"(value.w)
+ : "memory"
+ );
+ }
+
+ __device__ __forceinline__ unsigned int with_flag(unsigned int bits) {
+ return bits | 1u;
+ }
+
+ __device__ __forceinline__ unsigned int clear_flag(unsigned int bits) {
+ return bits & ~1u;
+ }
+
+ __device__ __forceinline__ bool flag_is_set(unsigned int bits) {
+ return (bits & 1u) != 0u;
+ }
+
+ __device__ __forceinline__ bool is_elected_in_warp(const int target_warp) {
+ const unsigned int warp_id = static_cast<unsigned int>(threadIdx.x >> 5);
+ const unsigned int uniform_warp_id = __shfl_sync(0xffffffffu, warp_id, 0);
+ return (uniform_warp_id == static_cast<unsigned int>(target_warp)) && ptx::elect_sync(0xffffffffu);
+ }
+
+ __global__ void prefixsum_kernel(const float* input, float* output, float* block_sums) {
+ const int block_idx = blockIdx.x;
+ const int tid = threadIdx.x;
+ const int lane = tid & 31;
+ const int warp = tid >> 5;
+
+ const int worker_blocks = 128;
+ const int coordinator_block = 128;
+ const int starter_warp = 0;
+ const int finalizer_warp = 1;
+ const int warp_size = 32;
+ const int values_per_thread = 4;
+ const int values_per_local = warp_size * values_per_thread; // 128 floats
+ const int values_per_block_round = LOCAL_STAGES * values_per_local;
+ const int slot_float4s = LOCAL_STAGES * warp_size;
+ const int round_span = worker_blocks * values_per_block_round;
+ const int rounds = (FIXED_N + round_span - 1) / round_span;
+ const unsigned int full_mask = 0xffffffffu;
+
+ if (block_idx < worker_blocks) {
+ extern __shared__ float4 stage_data[];
+ __shared__ uint64_t load_mbar[GLOBAL_STAGES];
+ __shared__ uint64_t done_mbar[GLOBAL_STAGES];
+
+ if (tid == 0) {
+ #pragma unroll
+ for (int s = 0; s < GLOBAL_STAGES; ++s) {
+ ptx::mbarrier_init(&load_mbar[s], 1);
+ ptx::mbarrier_init(&done_mbar[s], 1);
+ }
+ }
+ __syncthreads();
+
+ if (warp == finalizer_warp) {
+ const int warmup = (rounds < GLOBAL_STAGES) ? rounds : GLOBAL_STAGES;
+ #pragma unroll
+ for (int r = 0; r < warmup; ++r) {
+ const int slot = r;
+ const int src_base_idx = r * round_span + block_idx * values_per_block_round;
+ int valid_values = FIXED_N - src_base_idx;
+ if (valid_values < 0) valid_values = 0;
+ if (valid_values > values_per_block_round) valid_values = values_per_block_round;
+ const int valid_local_stages = valid_values / values_per_local;
+ const uint32_t valid_bytes = static_cast<uint32_t>(valid_local_stages * warp_size * sizeof(float4));
+ if (valid_local_stages > 0) {
+ if (is_elected_in_warp(finalizer_warp)) {
+ float4* dst = &stage_data[slot * slot_float4s];
+ const float4* src = reinterpret_cast<const float4*>(input + src_base_idx);
+ ptx::cp_async_bulk(
+ ptx::space_shared, ptx::space_global,
+ dst, src, valid_bytes, &load_mbar[slot]
+ );
+ (void)ptx::mbarrier_arrive_expect_tx(
+ ptx::sem_release, ptx::scope_cta, ptx::space_shared,
+ &load_mbar[slot], valid_bytes
+ );
+ }
+ } else {
+ if (is_elected_in_warp(finalizer_warp)) {
+ (void)ptx::mbarrier_arrive(
+ ptx::sem_release, ptx::scope_cta, ptx::space_shared, &load_mbar[slot]
+ );
+ }
+ }
+ }
+
+ for (int round = 0; round < rounds; ++round) {
+ const int slot = round % GLOBAL_STAGES;
+ const uint32_t phase = static_cast<uint32_t>((round / GLOBAL_STAGES) & 1);
+ while (!ptx::mbarrier_try_wait_parity(
+ ptx::sem_acquire, ptx::scope_cta, &done_mbar[slot], phase
+ )) {}
+
+ const float* mailbox_ptr = &block_sums[slot * worker_blocks + block_idx];
+ unsigned int bits = 0u;
+ do {
+ bits = __float_as_uint(ld_relaxed_gpu_f32(mailbox_ptr));
+ } while (flag_is_set(bits));
+ const float block_base = __uint_as_float(clear_flag(bits));
+
+ const int out_base_idx = round * round_span + block_idx * values_per_block_round;
+ int valid_values = FIXED_N - out_base_idx;
+ if (valid_values < 0) valid_values = 0;
+ if (valid_values > values_per_block_round) valid_values = values_per_block_round;
+ const int valid_local_stages = valid_values / values_per_local;
+
+ for (int ls = 0; ls < valid_local_stages; ++ls) {
+ const int smem_idx = slot * slot_float4s + ls * warp_size + lane;
+ float4 out4 = stage_data[smem_idx];
+ out4.x += block_base;
+ out4.y += block_base;
+ out4.z += block_base;
+ out4.w += block_base;
+
+ const int dst_idx = out_base_idx + ls * values_per_local + lane * values_per_thread;
+ reinterpret_cast<float4*>(output)[dst_idx / 4] = out4;
+ }
+
+ const int fetch_round = round + GLOBAL_STAGES;
+ if (fetch_round < rounds) {
+ const int src_base_idx = fetch_round * round_span + block_idx * values_per_block_round;
+ int next_valid_values = FIXED_N - src_base_idx;
+ if (next_valid_values < 0) next_valid_values = 0;
+ if (next_valid_values > values_per_block_round) next_valid_values = values_per_block_round;
+ const int next_valid_local_stages = next_valid_values / values_per_local;
+ const uint32_t next_valid_bytes =
+ static_cast<uint32_t>(next_valid_local_stages * warp_size * sizeof(float4));
+ if (next_valid_local_stages > 0) {
+ if (is_elected_in_warp(finalizer_warp)) {
+ float4* dst = &stage_data[slot * slot_float4s];
+ const float4* src = reinterpret_cast<const float4*>(input + src_base_idx);
+ ptx::cp_async_bulk(
+ ptx::space_shared, ptx::space_global,
+ dst, src, next_valid_bytes, &load_mbar[slot]
+ );
+ (void)ptx::mbarrier_arrive_expect_tx(
+ ptx::sem_release, ptx::scope_cta, ptx::space_shared,
+ &load_mbar[slot], next_valid_bytes
+ );
+ }
+ } else {
+ if (is_elected_in_warp(finalizer_warp)) {
+ (void)ptx::mbarrier_arrive(
+ ptx::sem_release, ptx::scope_cta, ptx::space_shared, &load_mbar[slot]
+ );
+ }
+ }
+ }
+ }
+ } else if (warp == starter_warp) {
+ for (int round = 0; round < rounds; ++round) {
+ const int slot = round % GLOBAL_STAGES;
+ const uint32_t phase = static_cast<uint32_t>((round / GLOBAL_STAGES) & 1);
+ while (!ptx::mbarrier_try_wait_parity(
+ ptx::sem_acquire, ptx::scope_cta, &load_mbar[slot], phase
+ )) {}
+
+ const int round_base_idx = round * round_span + block_idx * values_per_block_round;
+ int valid_values = FIXED_N - round_base_idx;
+ if (valid_values < 0) valid_values = 0;
+ if (valid_values > values_per_block_round) valid_values = values_per_block_round;
+ const int valid_local_stages = valid_values / values_per_local;
+ const int valid_float4 = valid_local_stages * warp_size;
+ int thread_items = valid_float4 - lane * LOCAL_STAGES;
+ if (thread_items < 0) thread_items = 0;
+ if (thread_items > LOCAL_STAGES) thread_items = LOCAL_STAGES;
+
+ const int slot_base = slot * slot_float4s;
+ const int thread_base = slot_base + lane * LOCAL_STAGES;
+
+ // Phase 1: each thread scans its contiguous LOCAL_STAGES chunk and stores partials.
+ float thread_running = 0.0f;
+ for (int i = 0; i < thread_items; ++i) {
+ const int smem_idx = thread_base + i;
+ const float4 in4 = stage_data[smem_idx];
+
+ const float p0 = in4.x;
+ const float p1 = p0 + in4.y;
+ const float p2 = p1 + in4.z;
+ const float p3 = p2 + in4.w;
+
+ float4 partial4;
+ partial4.x = p0 + thread_running;
+ partial4.y = p1 + thread_running;
+ partial4.z = p2 + thread_running;
+ partial4.w = p3 + thread_running;
+ stage_data[smem_idx] = partial4;
+ thread_running += p3;
+ }
+
+ // Phase 2: warp scan over per-thread totals to get thread-level exclusive carry.
+ float thread_inclusive = thread_running;
+ #pragma unroll
+ for (int offset = 1; offset < warp_size; offset <<= 1) {
+ const float y = __shfl_up_sync(full_mask, thread_inclusive, offset);
+ if (lane >= offset) {
+ thread_inclusive += y;
+ }
+ }
+ const float thread_exclusive = thread_inclusive - thread_running;
+
+ // Publish mailbox immediately after phase 2 once block total is known.
+ if (lane == warp_size - 1) {
+ const unsigned int tagged_bits = with_flag(__float_as_uint(thread_inclusive));
+ st_relaxed_gpu_f32(
+ &block_sums[slot * worker_blocks + block_idx],
+ __uint_as_float(tagged_bits)
+ );
+ }
+
+ // Phase 3: reload contiguous chunk and add thread carry.
+ for (int i = 0; i < thread_items; ++i) {
+ const int smem_idx = thread_base + i;
+ float4 out4 = stage_data[smem_idx];
+ out4.x += thread_exclusive;
+ out4.y += thread_exclusive;
+ out4.z += thread_exclusive;
+ out4.w += thread_exclusive;
+ stage_data[smem_idx] = out4;
+ }
+
+ if (lane == warp_size - 1) {
+ (void)ptx::mbarrier_arrive(
+ ptx::sem_release, ptx::scope_cta, ptx::space_shared, &done_mbar[slot]
+ );
+ }
+ }
+ }
+ return;
+ }
+
+ if (block_idx == coordinator_block) {
+ if (tid >= 32) {
+ return;
+ }
+
+ float running_sum = 0.0f;
+ for (int round = 0; round < rounds; ++round) {
+ const int slot = round % GLOBAL_STAGES;
+ const float* slot_base_ptr = &block_sums[slot * worker_blocks];
+
+ float4 tagged4;
+ unsigned int bits0 = 0u;
+ unsigned int bits1 = 0u;
+ unsigned int bits2 = 0u;
+ unsigned int bits3 = 0u;
+ bool all_ready = false;
+ do {
+ const float* slot_ptr = &slot_base_ptr[4 * lane];
+ tagged4 = ld_relaxed_gpu_v4_f32(slot_ptr);
+
+ bits0 = __float_as_uint(tagged4.x);
+ bits1 = __float_as_uint(tagged4.y);
+ bits2 = __float_as_uint(tagged4.z);
+ bits3 = __float_as_uint(tagged4.w);
+
+ const bool ready4 =
+ flag_is_set(bits0) && flag_is_set(bits1) &&
+ flag_is_set(bits2) && flag_is_set(bits3);
+ all_ready = __all_sync(full_mask, ready4);
+ } while (!all_ready);
+
+ const float s0 = __uint_as_float(clear_flag(bits0));
+ const float s1 = __uint_as_float(clear_flag(bits1));
+ const float s2 = __uint_as_float(clear_flag(bits2));
+ const float s3 = __uint_as_float(clear_flag(bits3));
+
+ const float lane_total = s0 + s1 + s2 + s3;
+ float warp_scan = lane_total;
+
+ #pragma unroll
+ for (int offset = 1; offset < 32; offset <<= 1) {
+ const float y = __shfl_up_sync(full_mask, warp_scan, offset);
+ if (lane >= offset) {
+ warp_scan += y;
+ }
+ }
+
+ const float lane_base = warp_scan - lane_total;
+ const float o0 = running_sum + lane_base;
+ const float o1 = o0 + s0;
+ const float o2 = o1 + s1;
+ const float o3 = o2 + s2;
+
+ float4 out4;
+ out4.x = __uint_as_float(clear_flag(__float_as_uint(o0)));
+ out4.y = __uint_as_float(clear_flag(__float_as_uint(o1)));
+ out4.z = __uint_as_float(clear_flag(__float_as_uint(o2)));
+ out4.w = __uint_as_float(clear_flag(__float_as_uint(o3)));
+ st_relaxed_gpu_v4_f32(&block_sums[slot * worker_blocks + 4 * lane], out4);
+
+ const float round_total = __shfl_sync(full_mask, warp_scan, 31);
+ running_sum += round_total;
+ }
+ }
+ }
+
+ torch::Tensor prefixsum_cuda(torch::Tensor input) {
+ const int size = input.numel();
+ auto output = torch::empty({size}, input.options());
+ auto block_sums = torch::zeros({128 * GLOBAL_STAGES}, input.options());
+ constexpr int threads_per_block = 64;
+ const size_t shared_bytes =
+ static_cast<size_t>(GLOBAL_STAGES) * LOCAL_STAGES * 32 * sizeof(float4);
+ cudaFuncSetAttribute(
+ prefixsum_kernel,
+ cudaFuncAttributeMaxDynamicSharedMemorySize,
+ static_cast<int>(shared_bytes)
+ );
+ cudaFuncSetAttribute(
+ prefixsum_kernel,
+ cudaFuncAttributePreferredSharedMemoryCarveout,
+ 100
+ );
+ prefixsum_kernel<<<129, threads_per_block, shared_bytes>>>(
+ input.data_ptr<float>(),
+ output.data_ptr<float>(),
+ block_sums.data_ptr<float>()
+ );
+ return output;
+ }
+ """
+
+
+ prefixsum_cpp_source = """
+ #include <torch/extension.h>
+
+ torch::Tensor prefixsum_cuda(torch::Tensor a);
+ """
+
+ prefixsum_module = load_inline(
+ name="prefixsum_cuda_dualwarp_tma",
+ cpp_sources=prefixsum_cpp_source,
+ cuda_sources=prefixsum_cuda_source,
+ functions=["prefixsum_cuda"],
+ verbose=True,
+ extra_cuda_cflags=[
+ "-O3",
+ f"-DGLOBAL_STAGES={GLOBAL_STAGES}",
+ f"-DLOCAL_STAGES={LOCAL_STAGES}",
+ f"-DFIXED_N={FIXED_N}",
+ "-gencode=arch=compute_90,code=sm_90",
+ ],
+ extra_ldflags=["-lcuda"],
+ )
+
+
def custom_kernel(data: input_t) -> output_t:
- data, output = data
- output = torch.cumsum(data, dim=0)
- return output
No newline at end of file
+ data = data[0]
+ if data.numel() != FIXED_N:
+ return torch.cumsum(data, dim=0)
+ return prefixsum_module.prefixsum_cuda(data)
scrolls · 441 diff lines total

Best evidence level for this revision: reported

JSON