Skip to content
KernelIndex
Search⌘K

submission 66191

achal · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-66191?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
13.5ms
#24 of 25
2025-10-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3ffe9c5dd1590efc65cb405f197412252cf5b01df62d58c64f61d7f25b76fe77
license declaredunknown
license concludedunknown
authorsachal
imported2026-08-15

Techniques

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

shared-memory__shared__ T segment[coarse_factor*block_dim];

Kernel source

submission.py842 lines
from utils import match_reference, DeterministicContext

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

from typing import Callable
from types import ModuleType

cuda_source = """
#define WORK_EFFICIENT 0

#ifndef CORE_TYPES_H
#define CORE_TYPES_H

#if defined(__clang__)
#define COMPILER_CLANG 1
#elif defined(_MSC_VER)
#define COMPILER_MSVC 1
#elif defined(__NVCC__)
#define COMPILER_NVCC
#else
#error "Compiler not supported"
#endif

#if defined(_M_X64) || defined(__x86_64__)
#define ARCH_X64 1
#elif defined(__aarch64__)
#define ARCH_ARM64 1
#elif defined(__wasm32__)
#define ARCH_WASM32 1
#else
#error "Architecture not supported"
#endif

#if defined(_WIN32)
#define PLATFORM_WINDOWS 1
#elif defined(__linux__)
#define PLATFORM_LINUX 1
#elif defined(__APPLE__)
#define PLATFORM_MACOS 1
#elif defined(__wasm32__)
#define PLATFORM_WASM 1
#else
#error "Platform not supported"
#endif

#define KiloBytes(x) (u64)(1024ull*(x))
#define MegaBytes(x) (u64)(1024ull*KiloBytes(x))
#define GigaBytes(x) (u64)(1024ull*MegaBytes(x))

#define Minimum(x, y) ((x) < (y) ? (x) : (y))
#define Maximum(x, y) ((x) > (y) ? (x) : (y))
#define ArrayCount(x) (sizeof(x)/sizeof((x)[0]))

#define CONCAT_IMPL(a, b) a##b
#define CONCAT(a, b) CONCAT_IMPL(a, b)

#if PLATFORM_WINDOWS
#define COBJMACROS
#define WIN32_LEAN_AND_MEAN
#include <windows.h>
#endif

#include <stdint.h>

// TODO(achal): Platform-specific asserts.
#if !ARCH_WASM32
#include <assert.h>
#define Assert assert
#else
#define Assert(...)
#endif

#define StaticAssert(cond, label) u8 static_assert_##label[(cond) ? (1) : (-1)];

typedef uint8_t u8;
typedef uint16_t u16;
typedef uint32_t u32;
typedef uint64_t u64;

typedef u8 b8;
typedef u32 b32;

typedef int16_t s16;
typedef int32_t s32;
typedef int64_t s64;

typedef float f32;
typedef double f64;

#endif // CORE_TYPES_H


#ifndef CORE_MEMORY_H
#define CORE_MEMORY_H

#if defined(ARCH_X64)
#define PageSize (KiloBytes(4))
#elif defined(ARCH_ARM64) || defined(ARCH_WASM32)
#define PageSize (KiloBytes(64))
#endif

typedef struct Arena Arena;

static Arena *g_scratch_arena = 0;

#if PLATFORM_WINDOWS
inline static u8 *OS_MemoryReserve(u64 size)
{
	u8 *memory = (u8 *)VirtualAlloc(0, size, MEM_RESERVE, PAGE_READWRITE);
	return memory;
}

inline static void OS_MemoryCommit(void *memory, u64 size)
{
	VirtualAlloc(memory, size, MEM_COMMIT, PAGE_READWRITE);
}

// NOTE(achal): My experiments show that calling VirtualAlloc twice - first with MEM_RESERVE and next with MEM_COMMIT - is much slower than calling it once with both the flags, hence the existence of this function.
inline static u8 *OS_MemoryReserveAndCommit(u64 size)
{
    u8 *memory = (u8 *)VirtualAlloc(0, size, MEM_RESERVE|MEM_COMMIT, PAGE_READWRITE);
    return memory;
}

static void OS_MemoryDecommit(void *memory, u64 size)
{
	int retval = VirtualFree(memory, size, MEM_DECOMMIT);
	Assert(retval);
}

static void OS_MemoryRelease(void *memory, u64 size)
{
	VirtualFree(memory, size, MEM_RELEASE);
}
#elif PLATFORM_LINUX || PLATFORM_MACOS
#include <sys/mman.h>
inline static u8 *OS_MemoryReserve(u64 size)
{
	u8 *memory = (u8 *)mmap(0, size, PROT_NONE, MAP_PRIVATE|MAP_ANON, -1, 0);
	Assert((void *)memory != MAP_FAILED);
	return memory;
}

static void OS_MemoryCommit(void *memory, u64 size)
{
	mprotect(memory, size, PROT_READ|PROT_WRITE);
}

inline static u8 *OS_MemoryReserveAndCommit(u64 size)
{
    u8 *result = (u8 *)mmap(0, size, PROT_READ|PROT_WRITE, MAP_PRIVATE|MAP_ANON, -1, 0);
	Assert((void *)result != MAP_FAILED);
    return result;
}

static void OS_MemoryDecommit(void *memory, u64 size)
{
	madvise(memory, size, MADV_DONTNEED);
	mprotect(memory, size, PROT_NONE);
}

static void OS_MemoryRelease(void *memory, u64 size)
{
	munmap(memory, size);
}
#endif

#if defined(ARCH_X64) || defined(ARCH_ARM64)
#include <string.h> // memset, memcpy

struct Arena
{
	u64 cap;
	u64 committed;
	u64 offset;
    
    // NOTE(achal): This points to the Arena struct itself which is followed
    // by the data it allocates i.e. the data really starts at
    // base + sizeof(Arena).
    //
	u8 *base;
};

// NOTE(achal):
// 1. `size` includes the size for the Arena struct itself.
// 2. `size` will get rounded up to next page boundary.
inline static Arena *InitArena(u64 size)
{
	u64 page_count = (size + PageSize - 1)/PageSize;
	size = page_count*PageSize;
    
	u8 *memory = OS_MemoryReserve(size);
	
	if (!memory) // @Investigate: How to handle errors here?
	{
		Assert(0);
		return 0;
	}
    
	u64 initial_commit_size = PageSize;
	Assert(initial_commit_size >= sizeof(Arena));
    
	OS_MemoryCommit(memory, initial_commit_size);
    
	Arena *arena = (Arena *)memory;
	arena->cap = size;
	arena->committed = initial_commit_size;
	arena->offset = sizeof(Arena);
	arena->base = memory;
    
	return arena;
}

inline static Arena *InitArenaDefault()
{
	u64 size = GigaBytes(8);
	Arena *result = InitArena(size);
	return result;
}

inline static u8 *GetArenaPos(Arena *arena)
{
    u8 *pos = arena->base + arena->offset;
    return pos;
}

inline static void ClearArena(Arena *arena)
{
	arena->offset = sizeof(Arena);
}

inline static void ReleaseArena(Arena *arena)
{
    // NOTE(achal): Since the address space for the entire arena was reserved by a single call to
    // VirtualAlloc, calling VirtualFree with zero will make sure that entire arena->cap worth of
    // address space gets released.
	OS_MemoryRelease(arena, 0);
}

static u8 *PushBytes(Arena *arena, u64 size)
{
	if (size == 0)
		return 0;
    
	u64 new_offset = arena->offset + size;
    
	if (new_offset > arena->cap)
	{
		Assert(!"Arena capacity reached!");
		return 0;
	}
    
	if (new_offset > arena->committed)
	{
		u64 size_to_commit = new_offset - arena->committed;
		u32 page_count = (u32)((size_to_commit + PageSize - 1) / PageSize);
		// TODO(achal): Handle failure.
		OS_MemoryCommit(arena->base + arena->committed, page_count*PageSize);
		arena->committed += page_count * PageSize;
	}
    
	u8 *result = arena->base + arena->offset;
	arena->offset = new_offset;
	return result;
}

static u8 *PushBytesZero(Arena *arena, u64 size)
{
	u8 *result = PushBytes(arena, size);
	memset(result, 0, size);
	return result;
}

static void PopBytesTo(Arena *arena, u64 offset)
{
	Assert(offset < arena->committed);
	u64 size_to_decommit = arena->committed - offset;
	u32 page_count = (u32)(size_to_decommit/PageSize);
	if (page_count)
	{
		arena->committed -= page_count * PageSize;
		OS_MemoryDecommit(arena->base + arena->committed, page_count*PageSize);
	}
    
	arena->offset = offset;
}

typedef struct
{
    Arena *arena;
    u64 offset;
} Scratch;

inline static Scratch ScratchBegin(Arena *arena)
{
    Scratch scratch = {arena, arena->offset};
    return scratch;
}

inline static void ScratchEnd(Scratch *scratch)
{
    scratch->arena->offset = scratch->offset;
}

static void MemoryCopy(u8 *dst, u8 *src, u64 size)
{
    memcpy(dst, src, size);
}
#elif defined(ARCH_WASM32)
// NOTE(achal):
// Making ArenaChunk fixed size means that when we get allocations that are larger than
// the size of the chunk we will have to allocate more than one ArenaChunk for a single
// allocation. Since, in general, ArenaChunks are not contiguous, we can't expect to
// allocation to be contiguous as well -- which is a guarantee Arena provides.
typedef struct ArenaChunk ArenaChunk;
struct ArenaChunk
{
    ArenaChunk *next;
    ArenaChunk *prev;
    
    u64 size;
    u64 offset;
    u8 *base;
};

static ArenaChunk *g_free_chunks_first = 0;
static ArenaChunk *g_free_chunks_last = 0;

struct Arena
{
    ArenaChunk *first_chunk;
    ArenaChunk *last_chunk;
    
    u64 size;
    u64 cap;
};

static u8 *PushBytes(Arena *arena, u64 size)
{
    if (size == 0)
		return 0;
    
    if (arena->size + size > arena->cap)
    {
        // TODO(achal): Assert(!"Arena capacity reached!");
        return 0;
    }
    
    ArenaChunk *chunk = 0;
    {
        chunk = arena->last_chunk;
        if (!chunk || (chunk->offset + size > chunk->size))
        {
            chunk = 0;
            // grab a new chunk
            for (ArenaChunk *c = g_free_chunks_first; c != g_free_chunks_last; c = c->next)
            {
                if (c->offset + size < c->size)
                {
                    chunk = c;
                    {
                        if (chunk == g_free_chunks_first)
                        {
                            g_free_chunks_first = chunk->next;
                        }
                        else if (chunk == g_free_chunks_last)
                        {
                            g_free_chunks_last = chunk->prev;
                        }
                        else
                        {
                            chunk->prev->next = chunk->next;
                        }
                    }
                    break;
                }
            }
            
            if (!chunk)
            {
                // NOTE(achal): This is to make sure that we don't allocate lots of small chunks. The choice
                // of PageSize is arbitrary
                u64 min_chunk_size = PageSize;
                u64 allocation_size = Maximum(size, min_chunk_size);
                
                u8 *memory = Platform_MemoryCommit(allocation_size);
                if (!memory)
                {
                    // TODO(achal): Handle failure?
                    Assert(0);
                }
                chunk = (ArenaChunk *)memory;
                chunk->prev = 0;
                chunk->next = 0;
                chunk->size = allocation_size;
                chunk->offset = sizeof(ArenaChunk);
                chunk->base = memory;
            }
            DoublyLinkedList_Push(arena->first_chunk, arena->last_chunk, next, prev, chunk);
        }
    }
    Assert(chunk);
    
    // allocate from the chunk
    u8 *result = chunk->base + chunk->offset;
    chunk->offset += size;
    arena->size += size;
    
    return result;
}

typedef struct
{
    Arena *arena;
    u64 size;
    ArenaChunk *chunk;
    u64 offset;
} Scratch;

inline static Scratch ScratchBegin(Arena *arena)
{
    Scratch scratch = {arena, arena->size, arena->last_chunk, arena->last_chunk->offset};
    return scratch;
}

inline static void ScratchEnd(Scratch *scratch)
{
    for (ArenaChunk *c = scratch->arena->last_chunk; c != scratch->chunk; c = c->prev)
    {
        c->offset = sizeof(ArenaChunk);
        scratch->arena->size -= c->size;
        DoublyLinkedList_Push(g_free_chunks_first, g_free_chunks_last, next, prev, c);
    }
    scratch->arena->last_chunk = scratch->chunk;
    scratch->arena->last_chunk->offset = scratch->offset;
    scratch->arena->size = scratch->size;
}

static void MemoryCopy(u8 *dst, u8 *src, u64 size)
{
    // @Optimization(Speed): Can we do better? For amd32 and arm32 we can just include
    // string.h, but 32-bit support right now is only for wasm and I don't know if there
    // is a good way for this without including all of standard library and making your
    // wasm module super bloated.
    for (u64 i = 0; i < size; ++i)
        *dst++ = *src++;
}
#endif

inline static Arena *GetScratchArena(u64 size)
{
    if (!g_scratch_arena)
        g_scratch_arena = InitArena(size);
    
    Assert(g_scratch_arena->cap >= size);
    return g_scratch_arena;
}

#define PushStruct(arena, type) (type *)PushBytes(arena, sizeof(type))
#define PushArray(arena, type, count) (type *)PushBytes(arena, count*sizeof(type))
#define PushArrayZero(arena, type, count) (type *)PushBytesZero(arena, count*sizeof(type))
#define PushStructZero(arena, type) (type *)PushBytesZero(arena, sizeof(type))

#endif // CORE_MEMORY_H

#define CUDACheck(expr) (expr)

template <typename T, bool inclusive, u32 block_dim, u32 coarse_factor>
__device__ T SegmentScan_KoggeStone(u64 count, T *array)
{
    __shared__ T segment[coarse_factor*block_dim];

    for (int i = 0; i < coarse_factor; ++i)
    {
        int smem_index = i*block_dim + threadIdx.x;
        int gmem_index = blockIdx.x*coarse_factor*block_dim + smem_index;
        if constexpr (inclusive)
        {
            if (gmem_index < count)
                segment[smem_index] = array[gmem_index];
            else
                segment[smem_index] = T(0);
        }
        else
        {
            if ((gmem_index >= count) || ((threadIdx.x == 0) && (i == 0)))
                segment[smem_index] = T(0);
            else
                segment[smem_index] = array[gmem_index - 1];
        }
    }
    __syncthreads();

    // step I: thread local sequential scan
    {
        u32 start = threadIdx.x*coarse_factor;
        for (u32 j = 1; j < coarse_factor; ++j)
            segment[start + j] += segment[start + (j - 1)];
    }
    __syncthreads();

    // step II: strided scan
    for (int stride = 1; stride < block_dim; stride *= 2)
    {
        T temp = T(0);
        if (threadIdx.x >= stride)
            temp = segment[(threadIdx.x - stride)*coarse_factor + (coarse_factor - 1)];
        __syncthreads();

        segment[threadIdx.x*coarse_factor + (coarse_factor - 1)] += temp;
        __syncthreads();
    }

    // step III: fixup
    {
        u32 start = threadIdx.x*coarse_factor;
        if (start > 0)
        {
            T prev = segment[start - 1];
            for (u32 i = 0; i < coarse_factor - 1; ++i)
                segment[start + i] += prev;
        }
    }
    __syncthreads();

    T segment_result = segment[(block_dim - 1)*coarse_factor + (coarse_factor - 1)];
    if constexpr (!inclusive)
    {
        u64 index = blockIdx.x*coarse_factor*block_dim + (block_dim - 1)*coarse_factor + (coarse_factor - 1);
        if (index > count - 1)
            index = count - 1;
        segment_result += array[index];
    }

    for (int i = 0; i < coarse_factor; ++i)
    {
        int smem_index = i*block_dim + threadIdx.x;
        int gmem_index = blockIdx.x*coarse_factor*block_dim + smem_index;
        if (gmem_index < count)
            array[gmem_index] = segment[smem_index];
    }

    return segment_result;
}

template <typename T, bool inclusive, u32 block_dim, u32 coarse_factor>
__device__ T SegmentScan_BrentKung(u64 count, T *array)
{
    __shared__ T segment[2*coarse_factor*block_dim];

    for (int i = 0; i < 2*coarse_factor; ++i)
    {
        int smem_index = i*block_dim + threadIdx.x;
        int gmem_index = blockIdx.x*2*coarse_factor*block_dim + smem_index;

        if constexpr (inclusive)
        {
            if (gmem_index < count)
                segment[smem_index] = array[gmem_index];
            else
                segment[smem_index] = T(0);
        }
        else
        {
            if ((gmem_index >= count) || ((threadIdx.x == 0) && (i == 0)))
                segment[smem_index] = T(0);
            else
                segment[smem_index] = array[gmem_index - 1];
        }
    }
    __syncthreads();

    // step I: thread local sequential scan
    {
        u32 thread_start = (threadIdx.x*2)*coarse_factor;
        for (u32 i = 0; i < 2; ++i)
        {
            u32 start = thread_start + i*coarse_factor;
            for (u32 j = 1; j < coarse_factor; ++j)
                segment[start + j] += segment[start + (j-1)];
        }
    }
    __syncthreads();

    // step II.I: reduction tree
    for (int stride = 1*coarse_factor; stride <= block_dim*coarse_factor; stride *= 2)
    {
        int index = (threadIdx.x + 1)*2*stride - 1;
        if (index < 2*block_dim*coarse_factor)
            segment[index] += segment[index - stride];
        __syncthreads();
    }

    T segment_result = segment[(block_dim - 1)*2*coarse_factor + (2*coarse_factor - 1)];
    if constexpr (!inclusive)
    {
        u64 index = blockIdx.x*2*coarse_factor*block_dim + (block_dim - 1)*2*coarse_factor + (2*coarse_factor - 1);
        if (index > count - 1)
            index = count - 1;
        segment_result += array[index];
    }

    // step II.II: downward
    for (int stride = ((2*block_dim)/4)*coarse_factor; stride >= 1*coarse_factor; stride /= 2)
    {
        int index = (threadIdx.x + 1)*2*stride - 1;
        if (index + stride < 2*block_dim*coarse_factor)
            segment[index + stride] += segment[index];
        __syncthreads();
    }

    // step III: fixup
    {
        u32 thread_start = (threadIdx.x*2)*coarse_factor;
        for (u32 i = 0; i < 2; ++i)
        {
            u32 start = thread_start + i*coarse_factor;
            if (start > 0)
            {
                T prev = segment[start - 1];
                for (u32 j = 0; j < coarse_factor - 1; ++j)
                    segment[start + j] += prev;
            }
        }
    }
    __syncthreads();

    for (int i = 0; i < 2*coarse_factor; ++i)
    {
        int smem_index = i*block_dim + threadIdx.x;
        int gmem_index = blockIdx.x*2*coarse_factor*block_dim + smem_index;
        if (gmem_index < count)
            array[gmem_index] = segment[smem_index];
    }

    return segment_result;
}

template <typename T, bool inclusive, u32 block_dim, u32 coarse_factor>
__global__ void Scan_Upsweep(u64 count, T *array, T *summary)
{
#if WORK_EFFICIENT
    T segment_result = SegmentScan_BrentKung<T, inclusive, block_dim, coarse_factor>(count, array);
#else
    T segment_result = SegmentScan_KoggeStone<T, inclusive, block_dim, coarse_factor>(count, array);
#endif
    if (summary && (threadIdx.x == block_dim - 1))
        summary[blockIdx.x] = segment_result;
}

template <typename T, bool inclusive, u32 block_dim, u32 coarse_factor>
__global__ void Scan_Downsweep(u64 count, T *array, T *summary)
{
#if WORK_EFFICIENT
    u8 elements_per_thread = 2*coarse_factor;
#else
    u8 elements_per_thread = coarse_factor;
#endif

    if (blockIdx.x > 0)
    {
        T segment_result;
        if constexpr (inclusive)
            segment_result = summary[blockIdx.x - 1];
        else
            segment_result = summary[blockIdx.x];

        for (u8 i = 0; i < elements_per_thread; ++i)
        {
            u64 index = blockIdx.x*elements_per_thread*block_dim + i*block_dim + threadIdx.x;
            if (index < count)
                array[index] += segment_result;
        }
    }
}

template <typename T>
struct Scan_PassInput
{
    T *d_array;
    T *d_summary;
    u64 element_count;
};

template <typename T>
struct Scan_Input
{
    T *d_scratch;
    u32 pass_input_count;
    Scan_PassInput<T> *pass_inputs;
};

template <typename T, u32 block_dim, u32 coarse_factor>
static Scan_Input<T> *Scan_CreateInput(Arena *arena, u64 count, T *d_array)
{
    Scan_Input<T> *input = PushStruct(arena, Scan_Input<T>);

#if WORK_EFFICIENT
    u32 elements_per_block = 2*coarse_factor*block_dim;
#else
    u32 elements_per_block = coarse_factor*block_dim;
#endif

    // Allocate scratch
    {
        u64 size = 0ull;

        // Add the number of elements for the first pass..
        u64 out_count = (count + elements_per_block - 1)/elements_per_block;
        size += out_count*sizeof(*input->d_scratch);

        // Add the number of elements for the second pass.. and we are done because we can just ping-pong.
        out_count = (out_count + elements_per_block - 1)/elements_per_block;
        size += out_count*sizeof(*input->d_scratch);

        CUDACheck(cudaMalloc(&input->d_scratch, size));
    }

    input->pass_input_count = (u32)ceilf(logf(count)/logf(elements_per_block));
    input->pass_inputs = PushArrayZero(arena, Scan_PassInput<T>, input->pass_input_count);

    int out_count = (count + elements_per_block - 1)/elements_per_block; // output of the first pass
    T *d_pingpong[] = {input->d_scratch, input->d_scratch + out_count};
    u64 element_count = count;

    for (u32 i = 0; i < input->pass_input_count; ++i)
    {
        Scan_PassInput<T> *pass_input = input->pass_inputs + i;
        pass_input->d_array = d_array;
        pass_input->d_summary = d_pingpong[i % 2];
        pass_input->element_count = element_count;

        int grid_dim = (element_count + elements_per_block - 1)/elements_per_block;
        
        d_array = pass_input->d_summary;
        element_count = grid_dim;
    }

    input->pass_inputs[input->pass_input_count - 1].d_summary = 0; // the "top" pass doesn't require a summary

    return input;
}

template <typename T>
static void Scan_DestroyInput(Scan_Input<T> *input)
{
    CUDACheck(cudaFree(input->d_scratch));
}

template <typename T, bool inclusive, u32 block_dim, u32 coarse_factor>
static void Scan(Scan_Input<T> *input, cudaStream_t stream)
{
    u32 k = input->pass_input_count;

#if WORK_EFFICIENT
    u32 elements_per_block = 2*coarse_factor*block_dim;
#else
    u32 elements_per_block = coarse_factor*block_dim;
#endif

    for (u32 pass = 0; pass < k; ++pass) // k upsweeps
    {
        Scan_PassInput<T> *pass_input = input->pass_inputs + pass;

        int grid_dim = (pass_input->element_count + elements_per_block - 1)/elements_per_block;
        CUDACheck((Scan_Upsweep<T, inclusive, block_dim, coarse_factor><<<grid_dim, block_dim, 0, stream>>>(pass_input->element_count, pass_input->d_array, pass_input->d_summary)));
    }

    for (s32 pass = k - 2; pass >= 0; --pass) // k - 1 downsweeps
    {
        Scan_PassInput<T> *pass_input = input->pass_inputs + pass;

        int grid_dim = (pass_input->element_count + elements_per_block - 1)/elements_per_block;
        CUDACheck((Scan_Downsweep<T, inclusive, block_dim, coarse_factor><<<grid_dim, block_dim, 0, stream>>>(pass_input->element_count, pass_input->d_array, pass_input->d_summary)));
    }
}

void Scan(torch::Tensor inout)
{
    u64 array_count = inout.numel();
    f64 *d_input = inout.data_ptr<f64>();

    enum
    {
        block_dim = 512,
        coarse_factor = 4,
    };

    {
        Scratch scratch = ScratchBegin(GetScratchArena(GigaBytes(8)));
        Scan_Input<f64> *input = Scan_CreateInput<f64, block_dim, coarse_factor>(scratch.arena, array_count, d_input);
        Scan<f64, true, block_dim, coarse_factor>(input, 0);
        Scan_DestroyInput(input);
        ScratchEnd(&scratch);
    }
}
"""

def load_cuda_prefixsumv2() -> ModuleType:
    prefixsumv2_module: ModuleType = None
    cpp_source = """
    #include <torch/extension.h>
    void Scan(torch::Tensor inout);
    """
    prefixsumv2_module = load_inline(
        name="Scan",
        cpp_sources=cpp_source,
        cuda_sources=cuda_source,
        functions=["Scan"],
        # NOTE(achal): Disable nvcc warning #177 - function declared but never referenced
        extra_cuda_cflags=["--diag-suppress", "177"],
        verbose=True)
    
    return prefixsumv2_module

prefixsumv2_module = load_cuda_prefixsumv2()

def custom_kernel_(data: input_t) -> output_t:
    with DeterministicContext():
        input, _ = data
        input_f64 = input.to(torch.float64)
        prefixsumv2_module.Scan(input_f64)
        output = input_f64.to(torch.float32)
        return output

custom_kernel = custom_kernel_

# def custom_kernel(data: input_t) -> output_t:
#     """
#     Reference implementation of inclusive prefix sum using PyTorch.
#     Args:
#         data: Input tensor to compute prefix sum on
#     Returns:
#         Tensor containing the inclusive prefix sum
#     """
#     data, output = data
#     output = torch.cumsum(data, dim=0)
#     return output
scrolls · 842 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON