Skip to content
KernelIndex
Search⌘K

submission 67230

shellsmile15795 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-vectorsum-v2-67230?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
Vector sum reductionsuite of 6 cases
NVIDIA B200
46.7µs
#15 of 88
2025-11-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:eeb003b19a3186396b68c8b1d14609ffebf61ce808e228c462e064397d0f61e6
license declaredunknown
license concludedunknown
authorsshellsmile15795
imported2026-08-15

Techniques

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

fp4const size_t value_bytes = requested_blocks; // one byte per NVFP4 payload (two nibbles per byte not packed yet)
shared-memory__shared__ float shared[Threads / kWarpSize];
vector-width = float4__device__ __forceinline__ float4 load_stream(const float4* addr) {

Kernel source

submission.py469 lines
import os
import shutil
from pathlib import Path

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


_CPP_SRC = r"""
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAStream.h>
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <mutex>

namespace {

// Persistent workspace for partial reductions.
uint8_t* workspace_ptr = nullptr;
size_t workspace_capacity = 0;
std::mutex workspace_mutex;

size_t nvfp4_workspace_bytes(size_t requested_blocks) {
    if (requested_blocks == 0) {
        return 0;
    }
    const size_t value_bytes = requested_blocks; // one byte per NVFP4 payload (two nibbles per byte not packed yet)
    const size_t scales_offset = (value_bytes + sizeof(float) - 1) / sizeof(float) * sizeof(float);
    const size_t scale_bytes = requested_blocks * sizeof(float); // per-block scale stored as fp32 for extended dynamic range
    const size_t total = scales_offset + scale_bytes;
    const size_t alignment = alignof(float);
    const size_t aligned_total = ((total + alignment - 1) / alignment) * alignment;
    return aligned_total;
}

uint8_t* ensure_workspace(size_t requested_blocks) {
    std::lock_guard<std::mutex> guard(workspace_mutex);
    if (requested_blocks == 0) {
        return nullptr;
    }
    const size_t required_bytes = nvfp4_workspace_bytes(requested_blocks);
    if (workspace_capacity < required_bytes) {
        if (workspace_ptr != nullptr) {
            C10_CUDA_CHECK(cudaFree(workspace_ptr));
        }
        C10_CUDA_CHECK(cudaMalloc(&workspace_ptr, required_bytes));
        workspace_capacity = required_bytes;
    }
    return workspace_ptr;
}

struct WorkspaceFinalizer {
    ~WorkspaceFinalizer() {
        if (workspace_ptr != nullptr) {
            cudaFree(workspace_ptr);
            workspace_ptr = nullptr;
            workspace_capacity = 0;
        }
    }
};

WorkspaceFinalizer workspace_finalizer;

} // namespace

// Forward declarations of CUDA kernels.
extern "C" void vector_sum_stage1_kernel(
    const float* input,
    uint8_t* workspace,
    int64_t elements,
    int blocks,
    int threads,
    cudaStream_t stream);

extern "C" void vector_sum_finalize_kernel(
    const uint8_t* workspace,
    float* output,
    int64_t blocks,
    int threads,
    cudaStream_t stream);

torch::Tensor vector_sum(torch::Tensor input, torch::Tensor output) {
    TORCH_CHECK(input.is_cuda(), "vector_sum: input must be on CUDA");
    TORCH_CHECK(output.is_cuda(), "vector_sum: output must be on CUDA");
    TORCH_CHECK(input.is_contiguous(), "vector_sum: input must be contiguous");
    TORCH_CHECK(output.is_contiguous(), "vector_sum: output must be contiguous");
    TORCH_CHECK(output.numel() == 1, "vector_sum: output must have a single element");
    TORCH_CHECK(input.scalar_type() == at::ScalarType::Float, "vector_sum: input must be float32");
    TORCH_CHECK(output.scalar_type() == at::ScalarType::Float, "vector_sum: output must be float32");

    const auto elements = input.numel();
    if (elements == 0) {
        output.zero_();
        return output;
    }

    constexpr int threads = 256;
    const int sm_count = at::cuda::getCurrentDeviceProperties()->multiProcessorCount;
    const int max_blocks = sm_count * 6;
    const int blocks = std::min<int>(
        static_cast<int>((elements + threads - 1) / threads),
        std::max(1, max_blocks)
    );

    uint8_t* workspace = ensure_workspace(static_cast<size_t>(blocks));
    TORCH_CHECK(workspace != nullptr, "vector_sum: failed to reserve workspace");

    auto stream = at::cuda::getCurrentCUDAStream();

    vector_sum_stage1_kernel(
        input.data_ptr<float>(),
        workspace,
        elements,
        blocks,
        threads,
        stream.stream()
    );

    constexpr int final_threads = 128;
    vector_sum_finalize_kernel(
        workspace,
        output.data_ptr<float>(),
        blocks,
        final_threads,
        stream.stream()
    );

    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return output;
}
"""


_CUDA_SRC = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <stdint.h>
#include <math.h>
#include <float.h>

namespace {

constexpr int kWarpSize = 32;
constexpr float kNvfp4ScaleMin = FLT_MIN; // avoid zero scale
constexpr float kNvfp4ScaleMax = FLT_MAX; // honour fp32 range
constexpr float kInvMantissa = 1.0f / 7.0f;
constexpr int kVectorUnroll = 3;

__device__ __forceinline__ float load_stream(const float* addr) {
#if __CUDA_ARCH__ >= 800
    float value;
    asm volatile("ld.global.cs.f32 %0, [%1];" : "=f"(value) : "l"(addr));
    return value;
#else
    return __ldg(addr);
#endif
}

__device__ __forceinline__ float4 load_stream(const float4* addr) {
    float4 value;
#if __CUDA_ARCH__ >= 800
    asm volatile("ld.global.cs.v4.f32 {%0, %1, %2, %3}, [%4];"
                 : "=f"(value.x), "=f"(value.y), "=f"(value.z), "=f"(value.w)
                 : "l"(addr));
#else
    value = __ldg(addr);
#endif
    return value;
}

__device__ __forceinline__ void prefetch_global_l2(const void* addr) {
#if __CUDA_ARCH__ >= 800
    asm volatile("prefetch.global.L2 [%0];" :: "l"(addr));
#endif
}

__inline__ __device__ float warp_reduce_sum(float value) {
    unsigned mask = __activemask();
    for (int offset = kWarpSize / 2; offset > 0; offset >>= 1) {
        value += __shfl_down_sync(mask, value, offset);
    }
    return value;
}

__inline__ __device__ double warp_reduce_sum(double value) {
    unsigned mask = __activemask();
    for (int offset = kWarpSize / 2; offset > 0; offset >>= 1) {
        value += __shfl_down_sync(mask, value, offset);
    }
    return value;
}

__forceinline__ __device__ size_t nvfp4_scale_offset(size_t block_count) {
    const size_t align = sizeof(float);
    return ((block_count + align - 1) / align) * align;
}

__forceinline__ __device__ void encode_nvfp4_value(
    float value,
    uint8_t& encoded_value,
    float& encoded_scale
) {
    if (!isfinite(value)) {
        value = 0.0f;
    }

    float abs_value = fabsf(value);
    if (abs_value == 0.0f) {
        encoded_value = static_cast<uint8_t>(8);
        encoded_scale = 0.0f;
        return;
    }

    float scale = abs_value * kInvMantissa;
    scale = fmaxf(scale, kNvfp4ScaleMin);
    scale = fminf(scale, kNvfp4ScaleMax);
    encoded_scale = scale;
    const float inv_scale = __frcp_rn(encoded_scale);
    float scaled = value * inv_scale;
    int quantized = __float2int_rn(scaled);

    if (quantized > 7) {
        quantized = 7;
    }
    if (quantized < -8) {
        quantized = -8;
    }

    encoded_value = static_cast<uint8_t>(quantized + 8);
}

__forceinline__ __device__ float decode_nvfp4_value(
    uint8_t encoded_value,
    float encoded_scale
) {
    const int quantized = static_cast<int>(encoded_value) - 8;
    if (quantized == 0) {
        return 0.0f;
    }
    const float scale = encoded_scale;
    if (scale == 0.0f) {
        return 0.0f;
    }
    return static_cast<float>(quantized) * scale;
}

template <int Threads>
__global__ void vector_sum_stage1(
    const float* __restrict__ input,
    uint8_t* __restrict__ workspace,
    int64_t elements
) {
    float thread_sum_f = 0.0f;
    const int64_t stride = static_cast<int64_t>(gridDim.x) * blockDim.x;

    const uintptr_t addr = reinterpret_cast<uintptr_t>(input);
    int64_t head = 0;
    if ((addr & 0xF) != 0) {
        const int misalignment = (addr & 0xF) >> 2; // measured in floats
        head = 4 - misalignment;
        if (head > elements) {
            head = elements;
        }
    }

    int64_t prefix_idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    while (prefix_idx < head) {
        float partial = load_stream(input + prefix_idx);
        thread_sum_f += partial;
        prefix_idx += stride;
    }

    const float* aligned_input = input + head;
    const int64_t aligned_elements = elements - head;
    const int64_t vec_elements = aligned_elements / 4;
    const float4* __restrict__ input4 = reinterpret_cast<const float4*>(aligned_input);
    int64_t vec_idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    const int64_t vec_stride = stride;
    const int64_t vec_unrolled_stride = vec_stride * kVectorUnroll;
    const int64_t prefetch_stride = vec_stride * kVectorUnroll * 2;

    while (vec_idx + (kVectorUnroll - 1) * vec_stride < vec_elements) {
        const int64_t pf_idx = vec_idx + prefetch_stride;
        if (pf_idx < vec_elements) {
            prefetch_global_l2(&input4[pf_idx]);
        }
#pragma unroll
        for (int u = 0; u < kVectorUnroll; ++u) {
            float4 v = load_stream(&input4[vec_idx + vec_stride * u]);
            thread_sum_f += (v.x + v.y) + (v.z + v.w);
        }
        vec_idx += vec_unrolled_stride;
    }

    while (vec_idx < vec_elements) {
        float4 v = load_stream(&input4[vec_idx]);
        thread_sum_f += (v.x + v.y) + (v.z + v.w);
        vec_idx += vec_stride;
    }

    const int64_t tail_start = head + vec_elements * 4;
    int64_t idx = tail_start + static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    while (idx < elements) {
        float partial = load_stream(input + idx);
        thread_sum_f += partial;
        idx += stride;
    }

    float thread_sum = warp_reduce_sum(thread_sum_f);

    __shared__ float shared[Threads / kWarpSize];
    if ((threadIdx.x & (kWarpSize - 1)) == 0) {
        shared[threadIdx.x / kWarpSize] = thread_sum;
    }
    __syncthreads();

    float block_sum_f = 0.0f;
    if (threadIdx.x < Threads / kWarpSize) {
        block_sum_f = shared[threadIdx.x];
        block_sum_f = warp_reduce_sum(block_sum_f);
    }

    if (threadIdx.x == 0) {
        const float block_sum = block_sum_f;
        const size_t block_count = static_cast<size_t>(gridDim.x);
        uint8_t* values = workspace;
        const size_t scales_offset = nvfp4_scale_offset(block_count);
        float* scales = reinterpret_cast<float*>(workspace + scales_offset);

        uint8_t encoded_value = 8;
        float encoded_scale = 0.0f;
        encode_nvfp4_value(block_sum, encoded_value, encoded_scale);
        values[blockIdx.x] = encoded_value;
        scales[blockIdx.x] = encoded_scale;
    }
}

template <int Threads>
__global__ void vector_sum_finalize(
    const uint8_t* __restrict__ workspace,
    float* __restrict__ output,
    int64_t blocks
) {
    const size_t block_count = static_cast<size_t>(blocks);
    const uint8_t* values = workspace;
    const size_t scales_offset = nvfp4_scale_offset(block_count);
    const float* scales = reinterpret_cast<const float*>(workspace + scales_offset);

    double thread_sum = 0.0;
    for (int64_t idx = threadIdx.x; idx < blocks; idx += blockDim.x) {
        const float decoded = decode_nvfp4_value(values[idx], scales[idx]);
        thread_sum += static_cast<double>(decoded);
    }

    thread_sum = warp_reduce_sum(thread_sum);

    __shared__ double shared[Threads / kWarpSize];
    if ((threadIdx.x & (kWarpSize - 1)) == 0) {
        shared[threadIdx.x / kWarpSize] = thread_sum;
    }
    __syncthreads();

    if (threadIdx.x < Threads / kWarpSize) {
        double block_sum = shared[threadIdx.x];
        block_sum = warp_reduce_sum(block_sum);
        if (threadIdx.x == 0) {
            output[0] = static_cast<float>(block_sum);
        }
    }
}

} // namespace

extern "C" void vector_sum_stage1_kernel(
    const float* input,
    uint8_t* workspace,
    int64_t elements,
    int blocks,
    int threads,
    cudaStream_t stream
) {
    dim3 grid(blocks);
    dim3 block(threads);
    switch (threads) {
        case 256:
            vector_sum_stage1<256><<<grid, block, 0, stream>>>(
                input,
                workspace,
                elements
            );
            break;
        case 512:
            vector_sum_stage1<512><<<grid, block, 0, stream>>>(
                input,
                workspace,
                elements
            );
            break;
        default:
            vector_sum_stage1<128><<<grid, dim3(128), 0, stream>>>(
                input,
                workspace,
                elements
            );
    }
}

extern "C" void vector_sum_finalize_kernel(
    const uint8_t* workspace,
    float* output,
    int64_t blocks,
    int threads,
    cudaStream_t stream
) {
    dim3 grid(1);
    dim3 block(threads);
    switch (threads) {
        case 128:
            vector_sum_finalize<128><<<grid, block, 0, stream>>>(
                workspace,
                output,
                blocks
            );
            break;
        default:
            vector_sum_finalize<64><<<grid, dim3(64), 0, stream>>>(
                workspace,
                output,
                blocks
            );
    }
}
"""


_CACHE_ROOTS = (
    Path("/root/.cache/torch_extension/py312_cu128/vectorsum_perf"),
    Path("/root/.cache/torch_extensions/py312_cu128/vectorsum_perf"),
)
for cache_path in _CACHE_ROOTS:
    if cache_path.exists():
        shutil.rmtree(cache_path, ignore_errors=True)


_MODULE = load_inline(
    name="vectorsum_perf",
    cpp_sources=_CPP_SRC,
    cuda_sources=_CUDA_SRC,
    functions=["vector_sum"],
    extra_cflags=["-O3", "-std=c++17"],
    extra_cuda_cflags=[
        "-O3",
        "--std=c++17",
        "-gencode=arch=compute_120,code=sm_120",
    ],
    verbose=False,
)


def _custom_kernel(data: input_t) -> output_t:
    input_tensor, output_tensor = data
    result = _MODULE.vector_sum(input_tensor, output_tensor)
    return result.view(())


custom_kernel = _custom_kernel
scrolls · 469 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 66881.

+ import os
+ import shutil
+ from pathlib import Path
+
import torch
from torch.utils.cpp_extension import load_inline
-
from task import input_t, output_t
- CUDA_SRC = r"""
+ _CPP_SRC = r"""
#include <ATen/cuda/CUDAContext.h>
+ #include <c10/cuda/CUDAStream.h>
#include <torch/extension.h>
#include <cuda_runtime.h>
- #include <cuda.h>
-
- #include <algorithm>
#include <cstdint>
+ #include <mutex>
namespace {
- constexpr int kWarpSize = 32;
- constexpr int kAsyncStages = 2;
- constexpr int kCopiesPerThread = 4;
+ // Persistent workspace for partial reductions.
+ uint8_t* workspace_ptr = nullptr;
+ size_t workspace_capacity = 0;
+ std::mutex workspace_mutex;
- __device__ inline double warp_reduce_sum(double value) {
- #pragma unroll
- for (int offset = kWarpSize / 2; offset > 0; offset >>= 1) {
- value += __shfl_down_sync(0xffffffff, value, offset);
+ size_t nvfp4_workspace_bytes(size_t requested_blocks) {
+ if (requested_blocks == 0) {
+ return 0;
}
- return value;
+ const size_t value_bytes = requested_blocks; // one byte per NVFP4 payload (two nibbles per byte not packed yet)
+ const size_t scales_offset = (value_bytes + sizeof(float) - 1) / sizeof(float) * sizeof(float);
+ const size_t scale_bytes = requested_blocks * sizeof(float); // per-block scale stored as fp32 for extended dynamic range
+ const size_t total = scales_offset + scale_bytes;
+ const size_t alignment = alignof(float);
+ const size_t aligned_total = ((total + alignment - 1) / alignment) * alignment;
+ return aligned_total;
}
- #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
- __device__ inline void cp_async_commit_group() { asm volatile("cp.async.commit_group;\n" ::); }
- __device__ inline void cp_async_wait_group_0() { asm volatile("cp.async.wait_group 0;\n" ::); }
- __device__ inline void cp_async_wait_all() { asm volatile("cp.async.wait_all;\n" ::); }
- #else
- __device__ inline void cp_async_commit_group() {}
- __device__ inline void cp_async_wait_group_0() {}
- __device__ inline void cp_async_wait_all() {}
- #endif
-
- __device__ inline bool async_copy_float4(float4* dst,
- const float4* base,
- const long long index,
- const long long vec_numel) {
- const bool valid = index < vec_numel;
- #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
- if (valid) {
- const unsigned int smem_addr = static_cast<unsigned int>(__cvta_generic_to_shared(dst));
- const unsigned long long gmem_addr = reinterpret_cast<unsigned long long>(base + index);
- asm volatile(
- "cp.async.cg.shared.global.L2::128B [%0], [%1], 16;\n"
- :: "r"(smem_addr), "l"(gmem_addr));
- } else {
- dst->x = 0.f;
- dst->y = 0.f;
- dst->z = 0.f;
- dst->w = 0.f;
+ uint8_t* ensure_workspace(size_t requested_blocks) {
+ std::lock_guard<std::mutex> guard(workspace_mutex);
+ if (requested_blocks == 0) {
+ return nullptr;
}
- #else
- if (valid) {
- *dst = base[index];
- } else {
- dst->x = 0.f;
- dst->y = 0.f;
- dst->z = 0.f;
- dst->w = 0.f;
+ const size_t required_bytes = nvfp4_workspace_bytes(requested_blocks);
+ if (workspace_capacity < required_bytes) {
+ if (workspace_ptr != nullptr) {
+ C10_CUDA_CHECK(cudaFree(workspace_ptr));
+ }
+ C10_CUDA_CHECK(cudaMalloc(&workspace_ptr, required_bytes));
+ workspace_capacity = required_bytes;
}
- #endif
- return valid;
+ return workspace_ptr;
}
- template <int Threads>
- __global__ void vector_sum_kernel_fallback(const float* __restrict__ input,
- double* __restrict__ accumulator,
- const long long numel,
- const int use_vectorized) {
- extern __shared__ double warp_sums[];
- float thread_sum_f = 0.0f;
+ struct WorkspaceFinalizer {
+ ~WorkspaceFinalizer() {
+ if (workspace_ptr != nullptr) {
+ cudaFree(workspace_ptr);
+ workspace_ptr = nullptr;
+ workspace_capacity = 0;
+ }
+ }
+ };
- const long long thread_linear =
- static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
- const long long stride =
- static_cast<long long>(blockDim.x) * gridDim.x;
+ WorkspaceFinalizer workspace_finalizer;
- if (use_vectorized) {
- const float4* __restrict__ input4 = reinterpret_cast<const float4*>(input);
- const long long vec_numel = numel >> 2;
- constexpr int kVectorUnroll = 8;
- float partials[kVectorUnroll] = {0.0f};
+ } // namespace
- long long index = thread_linear;
- const long long vec_stride = stride;
- const long long unroll_stride = vec_stride * kVectorUnroll;
- for (; index + static_cast<long long>(kVectorUnroll - 1) * vec_stride < vec_numel;
- index += unroll_stride) {
- #pragma unroll
- for (int unroll = 0; unroll < kVectorUnroll; ++unroll) {
- const long long load_index =
- index + static_cast<long long>(unroll) * vec_stride;
- float4 v;
- #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
- asm volatile("ld.global.cs.v4.f32 {%0, %1, %2, %3}, [%4];"
- : "=f"(v.x), "=f"(v.y), "=f"(v.z), "=f"(v.w)
- : "l"(input4 + load_index));
- #elif defined(__CUDA_ARCH__)
- v = __ldg(input4 + load_index);
- #else
- v = input4[load_index];
- #endif
- const float pair0 = v.x + v.y;
- const float pair1 = v.z + v.w;
- partials[unroll] += pair0 + pair1;
- }
- }
- for (; index < vec_numel; index += vec_stride) {
- float4 v;
- #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
- asm volatile("ld.global.cs.v4.f32 {%0, %1, %2, %3}, [%4];"
- : "=f"(v.x), "=f"(v.y), "=f"(v.z), "=f"(v.w)
- : "l"(input4 + index));
- #elif defined(__CUDA_ARCH__)
- v = __ldg(input4 + index);
- #else
- v = input4[index];
- #endif
- const float pair0 = v.x + v.y;
- const float pair1 = v.z + v.w;
- partials[0] += pair0 + pair1;
- }
+ // Forward declarations of CUDA kernels.
+ extern "C" void vector_sum_stage1_kernel(
+ const float* input,
+ uint8_t* workspace,
+ int64_t elements,
+ int blocks,
+ int threads,
+ cudaStream_t stream);
- const long long processed = vec_numel << 2;
- float tail_sum = 0.f;
- for (long long idx = processed + thread_linear; idx < numel; idx += stride) {
- tail_sum += input[idx];
- }
- float vector_sum = tail_sum;
- #pragma unroll
- for (int unroll = 0; unroll < kVectorUnroll; ++unroll) {
- vector_sum += partials[unroll];
- }
- thread_sum_f += vector_sum;
- } else {
- float accum0 = 0.0f;
- float accum1 = 0.0f;
- bool toggle = false;
- for (long long idx = thread_linear; idx < numel; idx += stride) {
- const float value = input[idx];
- if (!toggle) {
- accum0 += value;
- } else {
- accum1 += value;
- }
- toggle = !toggle;
- }
- thread_sum_f += accum0 + accum1;
- }
+ extern "C" void vector_sum_finalize_kernel(
+ const uint8_t* workspace,
+ float* output,
+ int64_t blocks,
+ int threads,
+ cudaStream_t stream);
- const int lane = threadIdx.x & (kWarpSize - 1);
- const int warp_id = threadIdx.x >> 5;
- const int warp_count = (Threads + kWarpSize - 1) / kWarpSize;
+ torch::Tensor vector_sum(torch::Tensor input, torch::Tensor output) {
+ TORCH_CHECK(input.is_cuda(), "vector_sum: input must be on CUDA");
+ TORCH_CHECK(output.is_cuda(), "vector_sum: output must be on CUDA");
+ TORCH_CHECK(input.is_contiguous(), "vector_sum: input must be contiguous");
+ TORCH_CHECK(output.is_contiguous(), "vector_sum: output must be contiguous");
+ TORCH_CHECK(output.numel() == 1, "vector_sum: output must have a single element");
+ TORCH_CHECK(input.scalar_type() == at::ScalarType::Float, "vector_sum: input must be float32");
+ TORCH_CHECK(output.scalar_type() == at::ScalarType::Float, "vector_sum: output must be float32");
- double thread_sum = static_cast<double>(thread_sum_f);
- double warp_sum = warp_reduce_sum(thread_sum);
- if (lane == 0) {
- warp_sums[warp_id] = warp_sum;
+ const auto elements = input.numel();
+ if (elements == 0) {
+ output.zero_();
+ return output;
}
- __syncthreads();
- if (warp_id == 0) {
- double block_sum = lane < warp_count ? warp_sums[lane] : 0.0;
- block_sum = warp_reduce_sum(block_sum);
- if (lane == 0) {
- atomicAdd(accumulator, block_sum);
- }
- }
- }
+ constexpr int threads = 256;
+ const int sm_count = at::cuda::getCurrentDeviceProperties()->multiProcessorCount;
+ const int max_blocks = sm_count * 6;
+ const int blocks = std::min<int>(
+ static_cast<int>((elements + threads - 1) / threads),
+ std::max(1, max_blocks)
+ );
- __global__ void vector_sum_kernel_async(const float* __restrict__ input,
- double* __restrict__ accumulator,
- const long long numel,
- const long long vec_numel) {
- extern __shared__ unsigned char shared_storage[];
- float4* tile_storage = reinterpret_cast<float4*>(shared_storage);
- const int warp_count = (blockDim.x + kWarpSize - 1) / kWarpSize;
- const int warp_stride = warp_count * kWarpSize;
- const int stage_area = warp_stride * kCopiesPerThread;
- double* warp_sums =
- reinterpret_cast<double*>(tile_storage + kAsyncStages * stage_area);
+ uint8_t* workspace = ensure_workspace(static_cast<size_t>(blocks));
+ TORCH_CHECK(workspace != nullptr, "vector_sum: failed to reserve workspace");
- const long long thread_linear =
- static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
- const long long stride =
- static_cast<long long>(blockDim.x) * gridDim.x;
- const float4* __restrict__ input4 = reinterpret_cast<const float4*>(input);
+ auto stream = at::cuda::getCurrentCUDAStream();
- const long long iteration_stride = stride * kCopiesPerThread;
- const long long total_iters =
- (vec_numel + iteration_stride - 1) / iteration_stride;
- float thread_acc[kCopiesPerThread] = {};
+ vector_sum_stage1_kernel(
+ input.data_ptr<float>(),
+ workspace,
+ elements,
+ blocks,
+ threads,
+ stream.stream()
+ );
- int stage = 0;
- const int lane = threadIdx.x & (kWarpSize - 1);
- const int warp_id = threadIdx.x >> 5;
+ constexpr int final_threads = 128;
+ vector_sum_finalize_kernel(
+ workspace,
+ output.data_ptr<float>(),
+ blocks,
+ final_threads,
+ stream.stream()
+ );
- if (total_iters > 0) {
- const long long initial_base = thread_linear;
- #pragma unroll
- for (int copy = 0; copy < kCopiesPerThread; ++copy) {
- const long long index = initial_base + static_cast<long long>(copy) * stride;
- float4* dst =
- tile_storage + stage * stage_area + copy * warp_stride +
- warp_id * kWarpSize + lane;
- async_copy_float4(dst, input4, index, vec_numel);
- }
- cp_async_commit_group();
+ C10_CUDA_KERNEL_LAUNCH_CHECK();
+ return output;
+ }
+ """
- for (long long iteration = 0; iteration < total_iters; ++iteration) {
- cp_async_wait_group_0();
- const float4* stage_base =
- tile_storage + stage * stage_area + warp_id * kWarpSize;
- const long long iteration_base =
- thread_linear + iteration * iteration_stride;
- #pragma unroll
- for (int copy = 0; copy < kCopiesPerThread; ++copy) {
- const long long vec_index =
- iteration_base + static_cast<long long>(copy) * stride;
- if (vec_index < vec_numel) {
- const float4 value =
- stage_base[copy * warp_stride + lane];
- // Keep per-copy accumulators to increase ILP for the reduction.
- const float pair0 = value.x + value.y;
- const float pair1 = value.z + value.w;
- const float partial = pair0 + pair1;
- thread_acc[copy] += partial;
- }
- }
+ _CUDA_SRC = r"""
+ #include <cuda.h>
+ #include <cuda_runtime.h>
+ #include <stdint.h>
+ #include <math.h>
+ #include <float.h>
- const long long next_base =
- thread_linear + (iteration + 1) * iteration_stride;
- const int next_stage = stage ^ 1;
+ namespace {
- if (iteration + 1 < total_iters) {
- #pragma unroll
- for (int copy = 0; copy < kCopiesPerThread; ++copy) {
- float4* next_dst =
- tile_storage + next_stage * stage_area +
- copy * warp_stride + warp_id * kWarpSize + lane;
- const long long next_index =
- next_base + static_cast<long long>(copy) * stride;
- async_copy_float4(next_dst, input4, next_index, vec_numel);
- }
- cp_async_commit_group();
- }
+ constexpr int kWarpSize = 32;
+ constexpr float kNvfp4ScaleMin = FLT_MIN; // avoid zero scale
+ constexpr float kNvfp4ScaleMax = FLT_MAX; // honour fp32 range
+ constexpr float kInvMantissa = 1.0f / 7.0f;
+ constexpr int kVectorUnroll = 3;
- stage = next_stage;
- }
- cp_async_wait_all();
+ __device__ __forceinline__ float load_stream(const float* addr) {
+ #if __CUDA_ARCH__ >= 800
+ float value;
+ asm volatile("ld.global.cs.f32 %0, [%1];" : "=f"(value) : "l"(addr));
+ return value;
+ #else
+ return __ldg(addr);
+ #endif
+ }
+
+ __device__ __forceinline__ float4 load_stream(const float4* addr) {
+ float4 value;
+ #if __CUDA_ARCH__ >= 800
+ asm volatile("ld.global.cs.v4.f32 {%0, %1, %2, %3}, [%4];"
+ : "=f"(value.x), "=f"(value.y), "=f"(value.z), "=f"(value.w)
+ : "l"(addr));
+ #else
+ value = __ldg(addr);
+ #endif
+ return value;
+ }
+
+ __device__ __forceinline__ void prefetch_global_l2(const void* addr) {
+ #if __CUDA_ARCH__ >= 800
+ asm volatile("prefetch.global.L2 [%0];" :: "l"(addr));
+ #endif
+ }
+
+ __inline__ __device__ float warp_reduce_sum(float value) {
+ unsigned mask = __activemask();
+ for (int offset = kWarpSize / 2; offset > 0; offset >>= 1) {
+ value += __shfl_down_sync(mask, value, offset);
}
+ return value;
+ }
- const long long processed = vec_numel << 2;
- float tail_sum = 0.f;
- for (long long idx = processed + thread_linear; idx < numel; idx += stride) {
- tail_sum += input[idx];
+ __inline__ __device__ double warp_reduce_sum(double value) {
+ unsigned mask = __activemask();
+ for (int offset = kWarpSize / 2; offset > 0; offset >>= 1) {
+ value += __shfl_down_sync(mask, value, offset);
}
- float thread_sum_f = tail_sum;
- #pragma unroll
- for (int copy = 0; copy < kCopiesPerThread; ++copy) {
- thread_sum_f += thread_acc[copy];
+ return value;
+ }
+
+ __forceinline__ __device__ size_t nvfp4_scale_offset(size_t block_count) {
+ const size_t align = sizeof(float);
+ return ((block_count + align - 1) / align) * align;
+ }
+
+ __forceinline__ __device__ void encode_nvfp4_value(
+ float value,
+ uint8_t& encoded_value,
+ float& encoded_scale
+ ) {
+ if (!isfinite(value)) {
+ value = 0.0f;
}
- double thread_sum = static_cast<double>(thread_sum_f);
- double warp_sum = warp_reduce_sum(thread_sum);
- if (lane == 0) {
- warp_sums[warp_id] = warp_sum;
+ float abs_value = fabsf(value);
+ if (abs_value == 0.0f) {
+ encoded_value = static_cast<uint8_t>(8);
+ encoded_scale = 0.0f;
+ return;
}
- __syncthreads();
- if (warp_id == 0) {
- const int warp_count_shared = (blockDim.x + kWarpSize - 1) / kWarpSize;
- double block_sum = lane < warp_count_shared ? warp_sums[lane] : 0.0;
- block_sum = warp_reduce_sum(block_sum);
- if (lane == 0) {
- atomicAdd(accumulator, block_sum);
- }
+ float scale = abs_value * kInvMantissa;
+ scale = fmaxf(scale, kNvfp4ScaleMin);
+ scale = fminf(scale, kNvfp4ScaleMax);
+ encoded_scale = scale;
+ const float inv_scale = __frcp_rn(encoded_scale);
+ float scaled = value * inv_scale;
+ int quantized = __float2int_rn(scaled);
+
+ if (quantized > 7) {
+ quantized = 7;
}
+ if (quantized < -8) {
+ quantized = -8;
+ }
+
+ encoded_value = static_cast<uint8_t>(quantized + 8);
}
- template <int Threads>
- void launch_fallback(const float* input,
- double* accumulator,
- const long long numel,
- const int use_vectorized,
- const int blocks,
- cudaStream_t stream) {
- const int warp_count = (Threads + kWarpSize - 1) / kWarpSize;
- const size_t shared_bytes = static_cast<size_t>(warp_count) * sizeof(double);
- vector_sum_kernel_fallback<Threads>
- <<<blocks, Threads, shared_bytes, stream>>>(
- input, accumulator, numel, use_vectorized);
+ __forceinline__ __device__ float decode_nvfp4_value(
+ uint8_t encoded_value,
+ float encoded_scale
+ ) {
+ const int quantized = static_cast<int>(encoded_value) - 8;
+ if (quantized == 0) {
+ return 0.0f;
+ }
+ const float scale = encoded_scale;
+ if (scale == 0.0f) {
+ return 0.0f;
+ }
+ return static_cast<float>(quantized) * scale;
}
- template <int Threads, int Unroll>
- __global__ void vector_sum_kernel_stream(const float* __restrict__ input,
- double* __restrict__ accumulator,
- const long long numel) {
- extern __shared__ double warp_sums[];
+ template <int Threads>
+ __global__ void vector_sum_stage1(
+ const float* __restrict__ input,
+ uint8_t* __restrict__ workspace,
+ int64_t elements
+ ) {
+ float thread_sum_f = 0.0f;
+ const int64_t stride = static_cast<int64_t>(gridDim.x) * blockDim.x;
- const long long thread_linear =
- static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
- const long long stride =
- static_cast<long long>(blockDim.x) * gridDim.x;
- const float4* __restrict__ input4 = reinterpret_cast<const float4*>(input);
- const long long vec_numel = numel >> 2;
+ const uintptr_t addr = reinterpret_cast<uintptr_t>(input);
+ int64_t head = 0;
+ if ((addr & 0xF) != 0) {
+ const int misalignment = (addr & 0xF) >> 2; // measured in floats
+ head = 4 - misalignment;
+ if (head > elements) {
+ head = elements;
+ }
+ }
- float thread_acc_even = 0.0f;
- float thread_acc_odd = 0.0f;
+ int64_t prefix_idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
+ while (prefix_idx < head) {
+ float partial = load_stream(input + prefix_idx);
+ thread_sum_f += partial;
+ prefix_idx += stride;
+ }
- long long index = thread_linear;
- const long long vec_stride = stride;
- const long long unroll_step =
- vec_stride * static_cast<long long>(Unroll);
- for (; index + static_cast<long long>(Unroll - 1) * vec_stride < vec_numel;
- index += unroll_step) {
+ const float* aligned_input = input + head;
+ const int64_t aligned_elements = elements - head;
+ const int64_t vec_elements = aligned_elements / 4;
+ const float4* __restrict__ input4 = reinterpret_cast<const float4*>(aligned_input);
+ int64_t vec_idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
+ const int64_t vec_stride = stride;
+ const int64_t vec_unrolled_stride = vec_stride * kVectorUnroll;
+ const int64_t prefetch_stride = vec_stride * kVectorUnroll * 2;
+
+ while (vec_idx + (kVectorUnroll - 1) * vec_stride < vec_elements) {
+ const int64_t pf_idx = vec_idx + prefetch_stride;
+ if (pf_idx < vec_elements) {
+ prefetch_global_l2(&input4[pf_idx]);
+ }
#pragma unroll
- for (int i = 0; i < Unroll; ++i) {
- const long long load_index =
- index + static_cast<long long>(i) * vec_stride;
- float4 v;
- #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
- asm volatile("ld.global.cs.v4.f32 {%0, %1, %2, %3}, [%4];"
- : "=f"(v.x), "=f"(v.y), "=f"(v.z), "=f"(v.w)
- : "l"(input4 + load_index));
- #elif defined(__CUDA_ARCH__)
- v = __ldg(input4 + load_index);
- #else
- v = input4[load_index];
- #endif
- const float pair0 = v.x + v.y;
- const float pair1 = v.z + v.w;
- const float partial = pair0 + pair1;
- if ((i & 1) == 0) {
- thread_acc_even += partial;
- } else {
- thread_acc_odd += partial;
- }
+ for (int u = 0; u < kVectorUnroll; ++u) {
+ float4 v = load_stream(&input4[vec_idx + vec_stride * u]);
+ thread_sum_f += (v.x + v.y) + (v.z + v.w);
}
+ vec_idx += vec_unrolled_stride;
}
- for (; index < vec_numel; index += vec_stride) {
- float4 v;
- #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
- asm volatile("ld.global.cs.v4.f32 {%0, %1, %2, %3}, [%4];"
- : "=f"(v.x), "=f"(v.y), "=f"(v.z), "=f"(v.w)
- : "l"(input4 + index));
- #elif defined(__CUDA_ARCH__)
- v = __ldg(input4 + index);
- #else
- v = input4[index];
- #endif
- const float pair0 = v.x + v.y;
- const float pair1 = v.z + v.w;
- thread_acc_even += pair0 + pair1;
+ while (vec_idx < vec_elements) {
+ float4 v = load_stream(&input4[vec_idx]);
+ thread_sum_f += (v.x + v.y) + (v.z + v.w);
+ vec_idx += vec_stride;
}
- const long long processed = vec_numel << 2;
- float tail_sum = 0.f;
- for (long long idx = processed + thread_linear; idx < numel; idx += stride) {
- tail_sum += input[idx];
+ const int64_t tail_start = head + vec_elements * 4;
+ int64_t idx = tail_start + static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
+ while (idx < elements) {
+ float partial = load_stream(input + idx);
+ thread_sum_f += partial;
+ idx += stride;
}
- float thread_sum_f = tail_sum;
- thread_sum_f += thread_acc_even + thread_acc_odd;
+ float thread_sum = warp_reduce_sum(thread_sum_f);
- double thread_sum = static_cast<double>(thread_sum_f);
- const int lane = threadIdx.x & (kWarpSize - 1);
- const int warp_id = threadIdx.x >> 5;
- const int warp_count = (Threads + kWarpSize - 1) / kWarpSize;
-
- double warp_sum = warp_reduce_sum(thread_sum);
- if (lane == 0) {
- warp_sums[warp_id] = warp_sum;
+ __shared__ float shared[Threads / kWarpSize];
+ if ((threadIdx.x & (kWarpSize - 1)) == 0) {
+ shared[threadIdx.x / kWarpSize] = thread_sum;
}
__syncthreads();
- if (warp_id == 0) {
- double block_sum = lane < warp_count ? warp_sums[lane] : 0.0;
- block_sum = warp_reduce_sum(block_sum);
- if (lane == 0) {
- atomicAdd(accumulator, block_sum);
- }
+ float block_sum_f = 0.0f;
+ if (threadIdx.x < Threads / kWarpSize) {
+ block_sum_f = shared[threadIdx.x];
+ block_sum_f = warp_reduce_sum(block_sum_f);
}
- }
- template <int Threads, int Unroll>
- void launch_stream(const float* input,
- double* accumulator,
- const long long numel,
- const int blocks,
- cudaStream_t stream) {
- const int warp_count = (Threads + kWarpSize - 1) / kWarpSize;
- const size_t shared_bytes =
- static_cast<size_t>(warp_count) * sizeof(double);
- vector_sum_kernel_stream<Threads, Unroll>
- <<<blocks, Threads, shared_bytes, stream>>>(input, accumulator, numel);
- }
+ if (threadIdx.x == 0) {
+ const float block_sum = block_sum_f;
+ const size_t block_count = static_cast<size_t>(gridDim.x);
+ uint8_t* values = workspace;
+ const size_t scales_offset = nvfp4_scale_offset(block_count);
+ float* scales = reinterpret_cast<float*>(workspace + scales_offset);
- } // namespace
-
- static inline int select_block_dim(const long long numel) {
- if (numel >= 1024) return 1024;
- if (numel >= 512) return 512;
- if (numel >= 256) return 256;
- if (numel >= 128) return 128;
- if (numel >= 64) return 64;
- if (numel >= 32) return 32;
- if (numel >= 16) return 16;
- if (numel >= 8) return 8;
- if (numel >= 4) return 4;
- return std::max<int>(static_cast<int>(numel), 1);
+ uint8_t encoded_value = 8;
+ float encoded_scale = 0.0f;
+ encode_nvfp4_value(block_sum, encoded_value, encoded_scale);
+ values[blockIdx.x] = encoded_value;
+ scales[blockIdx.x] = encoded_scale;
+ }
}
- void vector_sum_launcher(const at::Tensor& input,
- const at::Tensor& accumulator) {
- TORCH_CHECK(input.is_cuda(), "input must be on CUDA");
- TORCH_CHECK(accumulator.is_cuda(), "accumulator must be on CUDA");
- TORCH_CHECK(input.dtype() == torch::kFloat32, "input must be float32");
- TORCH_CHECK(accumulator.dtype() == torch::kFloat64, "accumulator must be float64");
- TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
- TORCH_CHECK(accumulator.is_contiguous(), "accumulator must be contiguous");
- TORCH_CHECK(accumulator.numel() == 1, "accumulator must have exactly one element");
+ template <int Threads>
+ __global__ void vector_sum_finalize(
+ const uint8_t* __restrict__ workspace,
+ float* __restrict__ output,
+ int64_t blocks
+ ) {
+ const size_t block_count = static_cast<size_t>(blocks);
+ const uint8_t* values = workspace;
+ const size_t scales_offset = nvfp4_scale_offset(block_count);
+ const float* scales = reinterpret_cast<const float*>(workspace + scales_offset);
- const long long numel = input.numel();
- if (numel == 0) {
- return;
+ double thread_sum = 0.0;
+ for (int64_t idx = threadIdx.x; idx < blocks; idx += blockDim.x) {
+ const float decoded = decode_nvfp4_value(values[idx], scales[idx]);
+ thread_sum += static_cast<double>(decoded);
}
- const bool vector_aligned =
- (numel >= 4) &&
- ((reinterpret_cast<std::uintptr_t>(input.data_ptr<float>()) & 0xF) == 0);
+ thread_sum = warp_reduce_sum(thread_sum);
- const auto* device_props = at::cuda::getCurrentDeviceProperties();
- const int supports_cp_async = device_props->major >= 8;
- const int sm_count = device_props->multiProcessorCount;
- const long long vec_numel = numel >> 2;
+ __shared__ double shared[Threads / kWarpSize];
+ if ((threadIdx.x & (kWarpSize - 1)) == 0) {
+ shared[threadIdx.x / kWarpSize] = thread_sum;
+ }
+ __syncthreads();
- int threads = select_block_dim(numel);
- const long long elements_per_thread = vector_aligned ? 4LL : 1LL;
- long long tentative_blocks =
- (numel + threads * elements_per_thread - 1) /
- (threads * elements_per_thread);
- tentative_blocks = std::max<long long>(tentative_blocks, 1LL);
- const int max_grid = 65535;
- int blocks = static_cast<int>(std::min<long long>(tentative_blocks, max_grid));
-
- const auto stream = at::cuda::getCurrentCUDAStream();
-
- constexpr int kStreamThreads = 640;
- constexpr int kStreamUnroll = 5;
- const bool use_stream = vector_aligned && numel >= 4096;
- if (use_stream) {
- threads = kStreamThreads;
- const long long stream_elements_per_thread =
- 4LL * static_cast<long long>(kStreamUnroll);
- long long stream_blocks =
- (numel + threads * stream_elements_per_thread - 1) /
- (threads * stream_elements_per_thread);
- stream_blocks = std::max<long long>(stream_blocks, 1LL);
- stream_blocks = std::min<long long>(stream_blocks, max_grid);
-
- const int warp_count_stream = (threads + kWarpSize - 1) / kWarpSize;
- const size_t shared_bytes_stream =
- static_cast<size_t>(warp_count_stream) * sizeof(double);
- int occupancy_blocks_per_sm = 0;
- if (cudaOccupancyMaxActiveBlocksPerMultiprocessor(
- &occupancy_blocks_per_sm,
- vector_sum_kernel_stream<kStreamThreads, kStreamUnroll>,
- threads,
- shared_bytes_stream) == cudaSuccess &&
- occupancy_blocks_per_sm > 0) {
- const long long occupancy_blocks =
- static_cast<long long>(sm_count) *
- static_cast<long long>(occupancy_blocks_per_sm);
- stream_blocks = std::min<long long>(
- stream_blocks,
- std::min<long long>(occupancy_blocks, max_grid));
- stream_blocks = std::max(stream_blocks, 1LL);
+ if (threadIdx.x < Threads / kWarpSize) {
+ double block_sum = shared[threadIdx.x];
+ block_sum = warp_reduce_sum(block_sum);
+ if (threadIdx.x == 0) {
+ output[0] = static_cast<float>(block_sum);
}
-
- launch_stream<kStreamThreads, kStreamUnroll>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(stream_blocks),
- stream);
- C10_CUDA_KERNEL_LAUNCH_CHECK();
- return;
}
+ }
- const bool use_async =
- supports_cp_async && vector_aligned && vec_numel > 0 && numel >= 1024;
+ } // namespace
- if (use_async) {
- threads = 768;
- const long long async_elements_per_thread =
- 4LL * static_cast<long long>(kCopiesPerThread);
- long long async_blocks =
- (numel + threads * async_elements_per_thread - 1) /
- (threads * async_elements_per_thread);
- async_blocks = std::max<long long>(async_blocks, 1LL);
- blocks = static_cast<int>(std::min<long long>(async_blocks, max_grid));
-
- const int warp_count = (threads + kWarpSize - 1) / kWarpSize;
- const size_t warp_stride = static_cast<size_t>(warp_count) * kWarpSize;
- const size_t stage_bytes =
- static_cast<size_t>(kAsyncStages) * warp_stride *
- static_cast<size_t>(kCopiesPerThread) * sizeof(float4);
- const size_t warp_bytes =
- static_cast<size_t>(warp_count) * sizeof(double);
- const size_t shared_bytes = stage_bytes + warp_bytes;
-
- int max_optin_smem = 0;
- C10_CUDA_CHECK(cudaDeviceGetAttribute(
- &max_optin_smem,
- cudaDevAttrMaxSharedMemoryPerBlockOptin,
- at::cuda::current_device()));
- TORCH_CHECK(
- shared_bytes <= static_cast<size_t>(max_optin_smem),
- "vector_sum_kernel_async requires ",
- shared_bytes,
- " bytes of shared memory, exceeding device limit of ",
- max_optin_smem);
- C10_CUDA_CHECK(cudaFuncSetAttribute(
- vector_sum_kernel_async,
- cudaFuncAttributeMaxDynamicSharedMemorySize,
- static_cast<int>(shared_bytes)));
-
- int occupancy_blocks_per_sm = 0;
- if (cudaOccupancyMaxActiveBlocksPerMultiprocessor(
- &occupancy_blocks_per_sm,
- vector_sum_kernel_async,
- threads,
- static_cast<size_t>(shared_bytes)) == cudaSuccess &&
- occupancy_blocks_per_sm > 0) {
- const long long occupancy_blocks =
- static_cast<long long>(sm_count) *
- static_cast<long long>(occupancy_blocks_per_sm);
- blocks = static_cast<int>(std::min<long long>(
- blocks,
- std::min<long long>(occupancy_blocks, max_grid)));
- blocks = std::max(blocks, 1);
- }
- vector_sum_kernel_async<<<blocks, threads, shared_bytes, stream>>>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- vec_numel);
- } else {
- switch (threads) {
- case 1024:
- launch_fallback<1024>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 512:
- launch_fallback<512>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 256:
- launch_fallback<256>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 128:
- launch_fallback<128>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 64:
- launch_fallback<64>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 32:
- launch_fallback<32>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 16:
- launch_fallback<16>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 8:
- launch_fallback<8>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 4:
- launch_fallback<4>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 3:
- launch_fallback<3>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 2:
- launch_fallback<2>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- case 1:
- launch_fallback<1>(
- input.data_ptr<float>(),
- accumulator.data_ptr<double>(),
- numel,
- static_cast<int>(vector_aligned),
- blocks,
- stream);
- break;
- default:
- TORCH_CHECK(false, "Unsupported block dimension");
- }
+ extern "C" void vector_sum_stage1_kernel(
+ const float* input,
+ uint8_t* workspace,
+ int64_t elements,
+ int blocks,
+ int threads,
+ cudaStream_t stream
+ ) {
+ dim3 grid(blocks);
+ dim3 block(threads);
+ switch (threads) {
+ case 256:
+ vector_sum_stage1<256><<<grid, block, 0, stream>>>(
+ input,
+ workspace,
+ elements
+ );
+ break;
+ case 512:
+ vector_sum_stage1<512><<<grid, block, 0, stream>>>(
+ input,
+ workspace,
+ elements
+ );
+ break;
+ default:
+ vector_sum_stage1<128><<<grid, dim3(128), 0, stream>>>(
+ input,
+ workspace,
+ elements
+ );
}
- C10_CUDA_KERNEL_LAUNCH_CHECK();
}
- PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
- m.def("vector_sum_launcher", &vector_sum_launcher, "vector sum reduction launcher");
+ extern "C" void vector_sum_finalize_kernel(
+ const uint8_t* workspace,
+ float* output,
+ int64_t blocks,
+ int threads,
+ cudaStream_t stream
+ ) {
+ dim3 grid(1);
+ dim3 block(threads);
+ switch (threads) {
+ case 128:
+ vector_sum_finalize<128><<<grid, block, 0, stream>>>(
+ workspace,
+ output,
+ blocks
+ );
+ break;
+ default:
+ vector_sum_finalize<64><<<grid, dim3(64), 0, stream>>>(
+ workspace,
+ output,
+ blocks
+ );
+ }
}
"""
- _module = load_inline(
- name="vectorsum_cuda_kernel",
- cpp_sources="",
- cuda_sources=CUDA_SRC,
- extra_cuda_cflags=["-O3", "--use_fast_math"],
+ _CACHE_ROOTS = (
+ Path("/root/.cache/torch_extension/py312_cu128/vectorsum_perf"),
+ Path("/root/.cache/torch_extensions/py312_cu128/vectorsum_perf"),
)
+ for cache_path in _CACHE_ROOTS:
+ if cache_path.exists():
+ shutil.rmtree(cache_path, ignore_errors=True)
- def custom_kernel(data: input_t) -> output_t:
- tensor, output = data # type: ignore[assignment]
+ _MODULE = load_inline(
+ name="vectorsum_perf",
+ cpp_sources=_CPP_SRC,
+ cuda_sources=_CUDA_SRC,
+ functions=["vector_sum"],
+ extra_cflags=["-O3", "-std=c++17"],
+ extra_cuda_cflags=[
+ "-O3",
+ "--std=c++17",
+ "-gencode=arch=compute_120,code=sm_120",
+ ],
+ verbose=False,
+ )
- if tensor.device.type != "cuda":
- raise RuntimeError("custom_kernel expects CUDA tensors")
- if not tensor.is_contiguous():
- tensor = tensor.contiguous()
+ def _custom_kernel(data: input_t) -> output_t:
+ input_tensor, output_tensor = data
+ result = _MODULE.vector_sum(input_tensor, output_tensor)
+ return result.view(())
- accumulator = torch.zeros(
- (),
- device=tensor.device,
- dtype=torch.float64,
- )
- _module.vector_sum_launcher(tensor, accumulator)
-
- result = accumulator.to(dtype=torch.float32)
- output.view(-1)[0] = result
- return result
+ custom_kernel = _custom_kernel
scrolls · 1082 diff lines total

Best evidence level for this revision: reported

JSON