Skip to content
KernelIndex
Search⌘K

submission 113167

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

template.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-113167?include=source"
interfacepython
Compatibility
measured onNVIDIA L4
declared hardwareNVIDIA L4
architecturessm_89
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
2D convolutionsuite of 5 cases
NVIDIA L4
1.26s
#19 of 21
2025-11-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f29f0aee62c3d1d7805bb6542fd7f119a67f76e9ce4ec7a7134149511ea94933
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15

Techniques

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

shared-memoryextern __shared__ float smem[];

Kernel source

template.py372 lines
import torch
from torch.utils.cpp_extension import load_inline

# --- 1. CUDA Kernel Source (核心逻辑不变,依然很快) ---
cuda_source = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <ATen/cuda/Exceptions.h>
#include <type_traits>

// Multi-output tiled convolution.
// Each block computes an output spatial tile for a small group of output channels,
// so the input patch in shared memory is reused across multiple channels and
// redundant global reads are reduced. A conservative shared-memory guard keeps
// launches within the 48 KB default limit; oversized cases fall back to a simple
// kernel that works for any shape.

namespace {

constexpr int kMaxSharedBytes = 48 * 1024;

template <int K, int TileX, int TileY, int OC_PER_BLOCK>
__global__ void conv2d_grouped_kernel(
    const float* __restrict__ input,
    const float* __restrict__ weight,
    float* __restrict__ output,
    int batch,
    int in_channels,
    int out_channels,
    int height,
    int width) {
    constexpr int tile_h = TileY + K - 1;
    constexpr int tile_w = TileX + K - 1;
    const int out_h = height - K + 1;
    const int out_w = width - K + 1;

    const int groups_per_batch = (out_channels + OC_PER_BLOCK - 1) / OC_PER_BLOCK;
    const int group_idx = blockIdx.z % groups_per_batch;
    const int b_idx = blockIdx.z / groups_per_batch;
    const int c_base = group_idx * OC_PER_BLOCK;
    const int c_out = c_base + threadIdx.z;

    // Spatial position for this thread.
    const int w_out = blockIdx.x * TileX + threadIdx.x;
    const int h_out = blockIdx.y * TileY + threadIdx.y;

    if (b_idx >= batch) {
        return;
    }

    extern __shared__ float smem[];
    float* smem_input = smem;  // [C, tile_h, tile_w]
    float* smem_kernel = smem_input + in_channels * tile_h * tile_w;  // [OC_PER_BLOCK, C, K, K]

    const int tid = (threadIdx.z * blockDim.y + threadIdx.y) * blockDim.x + threadIdx.x;
    const int threads = blockDim.x * blockDim.y * blockDim.z;

    const int input_elems = in_channels * tile_h * tile_w;
    const int kernel_elems_per_oc = in_channels * K * K;
    const int kernel_elems = kernel_elems_per_oc * OC_PER_BLOCK;

    // Stage input patch (shared across all output channels in the block).
    for (int idx = tid; idx < input_elems; idx += threads) {
        int tmp = idx;
        const int c = tmp / (tile_h * tile_w);
        tmp -= c * tile_h * tile_w;
        const int y = tmp / tile_w;
        const int x = tmp - y * tile_w;
        const int gy = blockIdx.y * TileY + y;
        const int gx = blockIdx.x * TileX + x;

        float val = 0.0f;
        if (gy < height && gx < width) {
            const int in_idx = ((b_idx * in_channels + c) * height + gy) * width + gx;
            val = input[in_idx];
        }
        smem_input[idx] = val;
    }

    // Stage kernels for the output-channel group handled by this block.
    for (int idx = tid; idx < kernel_elems; idx += threads) {
        int tmp = idx;
        const int oc_rel = tmp / kernel_elems_per_oc;
        tmp -= oc_rel * kernel_elems_per_oc;
        const int c = tmp / (K * K);
        tmp -= c * K * K;
        const int kh = tmp / K;
        const int kw = tmp - kh * K;

        float val = 0.0f;
        const int oc = c_base + oc_rel;
        if (oc < out_channels) {
            const int k_idx = ((oc * in_channels + c) * K + kh) * K + kw;
            val = weight[k_idx];
        }
        smem_kernel[idx] = val;
    }

    __syncthreads();

    if (w_out < out_w && h_out < out_h && c_out < out_channels) {
        float acc = 0.0f;
        const int out_base = threadIdx.y * tile_w + threadIdx.x;
        const float* kbase = smem_kernel + threadIdx.z * kernel_elems_per_oc;
#pragma unroll
        for (int c = 0; c < in_channels; ++c) {
            const float* tile = smem_input + c * tile_h * tile_w + out_base;
            const float* kptr = kbase + c * K * K;
#pragma unroll
            for (int kh = 0; kh < K; ++kh) {
#pragma unroll
                for (int kw = 0; kw < K; ++kw) {
                    acc += tile[kh * tile_w + kw] * kptr[kh * K + kw];
                }
            }
        }
        const int out_idx = ((b_idx * out_channels + c_out) * out_h + h_out) * out_w + w_out;
        output[out_idx] = acc;
    }
}

// Generic fallback when shared memory budget is exceeded or K is uncommon.
__global__ void conv2d_fallback_kernel(
    const float* __restrict__ input,
    const float* __restrict__ weight,
    float* __restrict__ output,
    int batch,
    int in_channels,
    int out_channels,
    int height,
    int width,
    int kernel_size) {
    const int w_out = blockIdx.x * blockDim.x + threadIdx.x;
    const int h_out = blockIdx.y * blockDim.y + threadIdx.y;
    const int b_idx = blockIdx.z / out_channels;
    const int c_out = blockIdx.z - b_idx * out_channels;

    const int out_h = height - kernel_size + 1;
    const int out_w = width - kernel_size + 1;
    if (w_out >= out_w || h_out >= out_h || b_idx >= batch) {
        return;
    }

    float acc = 0.0f;
    for (int c = 0; c < in_channels; ++c) {
        for (int kh = 0; kh < kernel_size; ++kh) {
            const int in_h = h_out + kh;
            for (int kw = 0; kw < kernel_size; ++kw) {
                const int in_w = w_out + kw;
                const int in_idx = ((b_idx * in_channels + c) * height + in_h) * width + in_w;
                const int k_idx = ((c_out * in_channels + c) * kernel_size + kh) * kernel_size + kw;
                acc += input[in_idx] * weight[k_idx];
            }
        }
    }
    const int out_idx = ((b_idx * out_channels + c_out) * out_h + h_out) * out_w + w_out;
    output[out_idx] = acc;
}

template <int K>
inline size_t shared_bytes_required(int in_channels, int tile_x, int tile_y, int oc_per_block) {
    const int tile_h = tile_y + K - 1;
    const int tile_w = tile_x + K - 1;
    const size_t input_bytes = static_cast<size_t>(in_channels) * tile_h * tile_w;
    const size_t kernel_bytes = static_cast<size_t>(oc_per_block) * in_channels * K * K;
    return (input_bytes + kernel_bytes) * sizeof(float);
}

template <int K, int TileX, int TileY, int OC_PER_BLOCK>
bool try_launch_cfg(
    const float* input,
    const float* weight,
    float* output,
    int batch,
    int in_channels,
    int out_channels,
    int height,
    int width,
    int out_h,
    int out_w) {
    constexpr int tile_x = TileX;
    constexpr int tile_y = TileY;
    constexpr int oc_per_block = OC_PER_BLOCK;
    const size_t shared = shared_bytes_required<K>(in_channels, tile_x, tile_y, oc_per_block);
    const int threads = tile_x * tile_y * oc_per_block;
    if (shared > kMaxSharedBytes || threads > 1024) {
        return false;
    }

    const int groups_per_batch = (out_channels + oc_per_block - 1) / oc_per_block;
    const dim3 block(tile_x, tile_y, oc_per_block);
    const dim3 grid((out_w + tile_x - 1) / tile_x,
                    (out_h + tile_y - 1) / tile_y,
                    batch * groups_per_batch);

    conv2d_grouped_kernel<K, TileX, TileY, OC_PER_BLOCK><<<grid, block, shared>>>(
        input, weight, output, batch, in_channels, out_channels, height, width);
    return true;
}

template <int K>
void launch_specialized(
    const float* input,
    const float* weight,
    float* output,
    int batch,
    int in_channels,
    int out_channels,
    int height,
    int width,
    int out_h,
    int out_w) {
    // Try higher OC grouping first to maximize input-tile reuse; fall back to smaller group sizes.
    if (try_launch_cfg<K, 16, 8, 4>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }
    if (try_launch_cfg<K, 8, 8, 4>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }
    if (try_launch_cfg<K, 16, 16, 2>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }
    if (try_launch_cfg<K, 16, 8, 2>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }
    if (try_launch_cfg<K, 8, 8, 2>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }
    if (try_launch_cfg<K, 16, 16, 1>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }
    if (try_launch_cfg<K, 16, 8, 1>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }
    if (try_launch_cfg<K, 8, 8, 1>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
        return;
    }

    // Shared memory would exceed cap; fall back.
    const dim3 fb_block(16, 16);
    const dim3 fb_grid((out_w + fb_block.x - 1) / fb_block.x,
                       (out_h + fb_block.y - 1) / fb_block.y,
                       batch * out_channels);
    conv2d_fallback_kernel<<<fb_grid, fb_block>>>(
        input, weight, output, batch, in_channels, out_channels, height, width, K);
}

}  // namespace

void conv2d_forward_wrapper(torch::Tensor input, torch::Tensor kernel, torch::Tensor output) {
    const int batch = input.size(0);
    const int in_channels = input.size(1);
    const int height = input.size(2);
    const int width = input.size(3);

    const int out_channels = kernel.size(0);
    const int kernel_size = kernel.size(2);

    const int out_h = height - kernel_size + 1;
    const int out_w = width - kernel_size + 1;

    auto launch = [&](auto KTag) {
        constexpr int K = decltype(KTag)::value;
        launch_specialized<K>(
            input.data_ptr<float>(),
            kernel.data_ptr<float>(),
            output.data_ptr<float>(),
            batch,
            in_channels,
            out_channels,
            height,
            width,
            out_h,
            out_w);
    };

    switch (kernel_size) {
        case 2:
            launch(std::integral_constant<int, 2>{});
            break;
        case 4:
            launch(std::integral_constant<int, 4>{});
            break;
        case 6:
            launch(std::integral_constant<int, 6>{});
            break;
        case 8:
            launch(std::integral_constant<int, 8>{});
            break;
        case 10:
            launch(std::integral_constant<int, 10>{});
            break;
        case 12:
            launch(std::integral_constant<int, 12>{});
            break;
        case 14:
            launch(std::integral_constant<int, 14>{});
            break;
        case 16:
            launch(std::integral_constant<int, 16>{});
            break;
        default: {
            const dim3 fb_block(16, 16);
            const dim3 fb_grid((out_w + fb_block.x - 1) / fb_block.x,
                               (out_h + fb_block.y - 1) / fb_block.y,
                               batch * out_channels);
            conv2d_fallback_kernel<<<fb_grid, fb_block>>>(
                input.data_ptr<float>(),
                kernel.data_ptr<float>(),
                output.data_ptr<float>(),
                batch,
                in_channels,
                out_channels,
                height,
                width,
                kernel_size);
            break;
        }
    }

    C10_CUDA_KERNEL_LAUNCH_CHECK();
}
"""

cpp_source = (
    "void conv2d_forward_wrapper("
    "    torch::Tensor input,"
    "    torch::Tensor kernel,"
    "    torch::Tensor output" # 对应修改签名
    ");"
)

# --- 2. 编译扩展 ---
conv2d_ext = load_inline(
    name="conv2d_v4_inplace",
    cpp_sources=[cpp_source],
    cuda_sources=[cuda_source],
    functions=["conv2d_forward_wrapper"],
    extra_cflags=["-std=c++17", "-O3"],
    extra_cuda_cflags=[
        "-O3",
        "--use_fast_math", 
        "-std=c++17",
        "--ptxas-options=-O3",
    ],
    verbose=False,
)

# --- 3. Python Wrapper (完全匹配官方 generate_input) ---
def custom_kernel(input_tuple):
    """
    Input: Tuple (input, kernel, output)
    Output: convolved result (returned specifically to satisfy any return checks)
    """
    # 1. 正确解包 3 个 Tensor
    input_tensor, kernel, output = input_tuple
    
    # 2. 确保连续性 (Contiguous)
    # 注意:如果 input/kernel 已经是 contiguous 的,这步操作开销几乎为 0
    if not input_tensor.is_contiguous():
        input_tensor = input_tensor.contiguous()
    if not kernel.is_contiguous():
        kernel = kernel.contiguous()
    # output 一般由 factory method 生成,默认是 contiguous 的,但为了安全也可以检查
    if not output.is_contiguous():
        output = output.contiguous()
    
    # 3. 调用 C++ (In-place 操作)
    conv2d_ext.conv2d_forward_wrapper(input_tensor, kernel, output)
    
    # 4. 必须返回 output,因为 leaderboard 通常会检查返回值
    return output
scrolls · 372 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 113115.

⋯ 7 unchanged lines
#include <ATen/cuda/Exceptions.h>
#include <type_traits>
- // Two-level tiled conv2d:
- // - Primary tile 16x8 outputs to increase reuse per block and SM occupancy.
- // - Fallback tile 8x8 when shared memory would exceed a conservative 48 KB cap.
- // - Final fallback kernel handles oversized cases or uncommon K.
+ // Multi-output tiled convolution.
+ // Each block computes an output spatial tile for a small group of output channels,
+ // so the input patch in shared memory is reused across multiple channels and
+ // redundant global reads are reduced. A conservative shared-memory guard keeps
+ // launches within the 48 KB default limit; oversized cases fall back to a simple
+ // kernel that works for any shape.
namespace {
constexpr int kMaxSharedBytes = 48 * 1024;
- template <int K, int TileX, int TileY>
- __global__ void conv2d_tiled_kernel(
+ template <int K, int TileX, int TileY, int OC_PER_BLOCK>
+ __global__ void conv2d_grouped_kernel(
const float* __restrict__ input,
const float* __restrict__ weight,
float* __restrict__ output,
⋯ 7 unchanged lines
const int out_h = height - K + 1;
const int out_w = width - K + 1;
+ const int groups_per_batch = (out_channels + OC_PER_BLOCK - 1) / OC_PER_BLOCK;
+ const int group_idx = blockIdx.z % groups_per_batch;
+ const int b_idx = blockIdx.z / groups_per_batch;
+ const int c_base = group_idx * OC_PER_BLOCK;
+ const int c_out = c_base + threadIdx.z;
+
+ // Spatial position for this thread.
const int w_out = blockIdx.x * TileX + threadIdx.x;
const int h_out = blockIdx.y * TileY + threadIdx.y;
- const int b_idx = blockIdx.z / out_channels;
- const int c_out = blockIdx.z - b_idx * out_channels;
if (b_idx >= batch) {
return;
}
extern __shared__ float smem[];
- float* smem_input = smem;
- float* smem_kernel = smem_input + in_channels * tile_h * tile_w;
+ float* smem_input = smem; // [C, tile_h, tile_w]
+ float* smem_kernel = smem_input + in_channels * tile_h * tile_w; // [OC_PER_BLOCK, C, K, K]
- const int tid = threadIdx.y * blockDim.x + threadIdx.x;
- const int threads = blockDim.x * blockDim.y;
+ const int tid = (threadIdx.z * blockDim.y + threadIdx.y) * blockDim.x + threadIdx.x;
+ const int threads = blockDim.x * blockDim.y * blockDim.z;
+
const int input_elems = in_channels * tile_h * tile_w;
- const int kernel_elems = in_channels * K * K;
+ const int kernel_elems_per_oc = in_channels * K * K;
+ const int kernel_elems = kernel_elems_per_oc * OC_PER_BLOCK;
- // Stage input patch covering the block output tile plus halo.
+ // Stage input patch (shared across all output channels in the block).
for (int idx = tid; idx < input_elems; idx += threads) {
int tmp = idx;
const int c = tmp / (tile_h * tile_w);
⋯ 11 unchanged lines
smem_input[idx] = val;
}
- // Stage filter weights for this output channel.
+ // Stage kernels for the output-channel group handled by this block.
for (int idx = tid; idx < kernel_elems; idx += threads) {
int tmp = idx;
+ const int oc_rel = tmp / kernel_elems_per_oc;
+ tmp -= oc_rel * kernel_elems_per_oc;
const int c = tmp / (K * K);
tmp -= c * K * K;
const int kh = tmp / K;
const int kw = tmp - kh * K;
- const int k_idx = ((c_out * in_channels + c) * K + kh) * K + kw;
- smem_kernel[idx] = weight[k_idx];
+
+ float val = 0.0f;
+ const int oc = c_base + oc_rel;
+ if (oc < out_channels) {
+ const int k_idx = ((oc * in_channels + c) * K + kh) * K + kw;
+ val = weight[k_idx];
+ }
+ smem_kernel[idx] = val;
}
__syncthreads();
- if (w_out < out_w && h_out < out_h) {
+ if (w_out < out_w && h_out < out_h && c_out < out_channels) {
float acc = 0.0f;
- const int out_base_x = threadIdx.x;
- const int out_base_y = threadIdx.y;
+ const int out_base = threadIdx.y * tile_w + threadIdx.x;
+ const float* kbase = smem_kernel + threadIdx.z * kernel_elems_per_oc;
+ #pragma unroll
for (int c = 0; c < in_channels; ++c) {
- const float* tile = smem_input + c * tile_h * tile_w + out_base_y * tile_w + out_base_x;
- const float* kptr = smem_kernel + c * K * K;
+ const float* tile = smem_input + c * tile_h * tile_w + out_base;
+ const float* kptr = kbase + c * K * K;
#pragma unroll
for (int kh = 0; kh < K; ++kh) {
#pragma unroll
⋯ 45 unchanged lines
output[out_idx] = acc;
}
- inline size_t shared_bytes_required(int in_channels, int kernel_size, int tile_x, int tile_y) {
- const int tile_h = tile_y + kernel_size - 1;
- const int tile_w = tile_x + kernel_size - 1;
- return static_cast<size_t>(in_channels) *
- static_cast<size_t>(tile_h * tile_w + kernel_size * kernel_size) * sizeof(float);
+ template <int K>
+ inline size_t shared_bytes_required(int in_channels, int tile_x, int tile_y, int oc_per_block) {
+ const int tile_h = tile_y + K - 1;
+ const int tile_w = tile_x + K - 1;
+ const size_t input_bytes = static_cast<size_t>(in_channels) * tile_h * tile_w;
+ const size_t kernel_bytes = static_cast<size_t>(oc_per_block) * in_channels * K * K;
+ return (input_bytes + kernel_bytes) * sizeof(float);
}
- template <int K>
- void launch_specialized(
+ template <int K, int TileX, int TileY, int OC_PER_BLOCK>
+ bool try_launch_cfg(
const float* input,
const float* weight,
float* output,
⋯ 4 unchanged lines
int width,
int out_h,
int out_w) {
- constexpr int fast_tile_x = 16;
- constexpr int fast_tile_y = 8;
- constexpr int compat_tile_x = 8;
- constexpr int compat_tile_y = 8;
+ constexpr int tile_x = TileX;
+ constexpr int tile_y = TileY;
+ constexpr int oc_per_block = OC_PER_BLOCK;
+ const size_t shared = shared_bytes_required<K>(in_channels, tile_x, tile_y, oc_per_block);
+ const int threads = tile_x * tile_y * oc_per_block;
+ if (shared > kMaxSharedBytes || threads > 1024) {
+ return false;
+ }
- const size_t fast_shared = shared_bytes_required(in_channels, K, fast_tile_x, fast_tile_y);
- const size_t compat_shared = shared_bytes_required(in_channels, K, compat_tile_x, compat_tile_y);
+ const int groups_per_batch = (out_channels + oc_per_block - 1) / oc_per_block;
+ const dim3 block(tile_x, tile_y, oc_per_block);
+ const dim3 grid((out_w + tile_x - 1) / tile_x,
+ (out_h + tile_y - 1) / tile_y,
+ batch * groups_per_batch);
- if (fast_shared <= kMaxSharedBytes) {
- const dim3 block(fast_tile_x, fast_tile_y);
- const dim3 grid((out_w + fast_tile_x - 1) / fast_tile_x,
- (out_h + fast_tile_y - 1) / fast_tile_y,
- batch * out_channels);
- conv2d_tiled_kernel<K, fast_tile_x, fast_tile_y><<<grid, block, fast_shared>>>(
- input, weight, output, batch, in_channels, out_channels, height, width);
+ conv2d_grouped_kernel<K, TileX, TileY, OC_PER_BLOCK><<<grid, block, shared>>>(
+ input, weight, output, batch, in_channels, out_channels, height, width);
+ return true;
+ }
+
+ template <int K>
+ void launch_specialized(
+ const float* input,
+ const float* weight,
+ float* output,
+ int batch,
+ int in_channels,
+ int out_channels,
+ int height,
+ int width,
+ int out_h,
+ int out_w) {
+ // Try higher OC grouping first to maximize input-tile reuse; fall back to smaller group sizes.
+ if (try_launch_cfg<K, 16, 8, 4>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
return;
}
-
- if (compat_shared <= kMaxSharedBytes) {
- const dim3 block(compat_tile_x, compat_tile_y);
- const dim3 grid((out_w + compat_tile_x - 1) / compat_tile_x,
- (out_h + compat_tile_y - 1) / compat_tile_y,
- batch * out_channels);
- conv2d_tiled_kernel<K, compat_tile_x, compat_tile_y><<<grid, block, compat_shared>>>(
- input, weight, output, batch, in_channels, out_channels, height, width);
+ if (try_launch_cfg<K, 8, 8, 4>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
return;
}
+ if (try_launch_cfg<K, 16, 16, 2>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
+ return;
+ }
+ if (try_launch_cfg<K, 16, 8, 2>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
+ return;
+ }
+ if (try_launch_cfg<K, 8, 8, 2>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
+ return;
+ }
+ if (try_launch_cfg<K, 16, 16, 1>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
+ return;
+ }
+ if (try_launch_cfg<K, 16, 8, 1>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
+ return;
+ }
+ if (try_launch_cfg<K, 8, 8, 1>(input, weight, output, batch, in_channels, out_channels, height, width, out_h, out_w)) {
+ return;
+ }
// Shared memory would exceed cap; fall back.
const dim3 fb_block(16, 16);
scrolls · 221 diff lines total

Best evidence level for this revision: reported

JSON