submission 115574
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 308 lines, June 9 Researcher Reciprocity License v1.0.
template.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-115574?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:c0435e487cad5b0f4ebdbf8f7dbc79e54c84cc6b9b6fcdc9ca494f01080dbd0c
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
alignas(16) __shared__ float k_smem[C][K][K][OC];vector-width = float4
const float4 w = reinterpret_cast<const float4*>(wptr)[oc / 4];Kernel source
template.py308 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 <algorithm>
// 辅助函数:float4 累加
template <int OC>
__device__ __forceinline__ void accumulate_oc(const float* wptr, float v, float* acc) {
if constexpr (OC % 4 == 0) {
#pragma unroll
for (int oc = 0; oc < OC; oc += 4) {
const float4 w = reinterpret_cast<const float4*>(wptr)[oc / 4];
acc[oc + 0] = fmaf(v, w.x, acc[oc + 0]);
acc[oc + 1] = fmaf(v, w.y, acc[oc + 1]);
acc[oc + 2] = fmaf(v, w.z, acc[oc + 2]);
acc[oc + 3] = fmaf(v, w.w, acc[oc + 3]);
}
} else {
#pragma unroll
for (int oc = 0; oc < OC; ++oc) {
acc[oc] = fmaf(v, wptr[oc], acc[oc]);
}
}
}
namespace {
template <int TW, int TH>
inline dim3 make_grid(int out_w, int out_h, int batch, int channel_chunk) {
const int grid_x = (out_w + TW - 1) / TW;
const int grid_y = (out_h + TH - 1) / TH;
const int grid_z = batch * channel_chunk;
return dim3(grid_x, grid_y, grid_z);
}
// ---- 专门化多通道静态共享内存核 (保持不变,这部分逻辑是好的) ----
template <int K, int C, int OC, int TW, int TH>
__launch_bounds__(TW * TH, 2)
__global__ void conv2d_multi_static(
const float* __restrict__ input,
const float* __restrict__ weight,
float* __restrict__ output,
int batch,
int out_channels,
int height,
int width) {
// ... 保持你的原始逻辑 ...
// 为了节省篇幅,这里省略重复代码,但编译时需要包含你的原始 conv2d_multi_static 实现
// 请确保这里包含你原代码中该函数的完整内容
static_assert(TW * TH <= 1024, "block too large");
constexpr int PATCH_W = TW + K - 1;
constexpr int PATCH_H = TH + K - 1;
constexpr int PATCH_STRIDE = (PATCH_W >= 15) ? PATCH_W : (PATCH_W + 1);
const int out_w = width - K + 1;
const int out_h = height - K + 1;
const int tile_x = blockIdx.x * TW;
const int tile_y = blockIdx.y * TH;
if (tile_x >= out_w || tile_y >= out_h) return;
const int channel_chunks = (out_channels + OC - 1) / OC;
const int bc = static_cast<int>(blockIdx.z);
const int b_idx = bc / channel_chunks;
const int chunk_idx = bc - b_idx * channel_chunks;
const int c_start = chunk_idx * OC;
if (b_idx >= batch || c_start >= out_channels) return;
alignas(16) __shared__ float k_smem[C][K][K][OC];
alignas(16) __shared__ float in_smem[C][PATCH_H][PATCH_STRIDE];
const int tid = threadIdx.y * blockDim.x + threadIdx.x;
const int blk_threads = blockDim.x * blockDim.y;
const int total_k = OC * C * K * K;
for (int idx = tid; idx < total_k; idx += blk_threads) {
const int tmp = idx / OC;
const int oc = idx - tmp * OC;
const int ci = tmp / (K * K);
const int rem0 = tmp - ci * K * K;
const int kh = rem0 / K;
const int kw = rem0 - kh * K;
const int g_oc = c_start + oc;
if (g_oc < out_channels) {
const int k_idx = ((g_oc * C + ci) * K + kh) * K + kw;
k_smem[ci][kh][kw][oc] = weight[k_idx];
} else {
k_smem[ci][kh][kw][oc] = 0.0f;
}
}
const int patch_elems = C * PATCH_H * PATCH_W;
for (int idx = tid; idx < patch_elems; idx += blk_threads) {
const int ci = idx / (PATCH_H * PATCH_W);
const int rem0 = idx - ci * PATCH_H * PATCH_W;
const int py = rem0 / PATCH_W;
const int px = rem0 - py * PATCH_W;
const int gx = tile_x + px;
const int gy = tile_y + py;
float val = 0.f;
if (gx < width && gy < height && b_idx < batch) {
const int in_idx = ((b_idx * C + ci) * height + gy) * width + gx;
val = input[in_idx];
}
in_smem[ci][py][px] = val;
}
__syncthreads();
const int out_x = tile_x + threadIdx.x;
const int out_y = tile_y + threadIdx.y;
if (out_x < out_w && out_y < out_h) {
float acc[OC];
#pragma unroll
for (int oc = 0; oc < OC; ++oc) acc[oc] = 0.f;
#pragma unroll
for (int ci = 0; ci < C; ++ci) {
const float* in_ptr = &in_smem[ci][threadIdx.y][threadIdx.x];
#pragma unroll
for (int kh = 0; kh < K; ++kh) {
#pragma unroll
for (int kw = 0; kw < K; ++kw) {
const float v = in_ptr[kh * PATCH_STRIDE + kw];
const float* wptr = &k_smem[ci][kh][kw][0];
accumulate_oc<OC>(wptr, v, acc);
}
}
}
#pragma unroll
for (int oc = 0; oc < OC; ++oc) {
const int g_oc = c_start + oc;
if (g_oc < out_channels) {
const int out_idx = ((b_idx * out_channels + g_oc) * out_h + out_y) * out_w + out_x;
output[out_idx] = acc[oc];
}
}
}
}
// ---- 【修复版】通用 fallback:Global Memory 版本 ----
// 不使用 Shared Memory,避免因 channel 过多导致 Invalid Argument
__global__ void conv2d_generic_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) {
// 每个 block 处理 16x16 的输出区域
const int out_w = width - kernel_size + 1;
const int out_h = height - kernel_size + 1;
const int col = blockIdx.x * blockDim.x + threadIdx.x;
const int row = blockIdx.y * blockDim.y + threadIdx.y;
// Z 轴映射到 Batch 和 Out Channel
const int bc = blockIdx.z;
const int b = bc / out_channels;
const int c_out = bc % out_channels; // 修正逻辑:bc = b * out_c + c_out
if (col >= out_w || row >= out_h || b >= batch) return;
float acc = 0.0f;
// 直接遍历输入通道,利用 L1 Cache 进行缓存
for (int c = 0; c < in_channels; ++c) {
for (int kh = 0; kh < kernel_size; ++kh) {
int in_h = row + kh;
for (int kw = 0; kw < kernel_size; ++kw) {
int in_w = col + kw;
// 计算全局索引
int in_idx = ((b * in_channels + c) * height + in_h) * width + in_w;
int k_idx = ((c_out * in_channels + c) * kernel_size + kh) * kernel_size + kw;
// 乘累加
acc += input[in_idx] * weight[k_idx];
}
}
}
int out_idx = ((b * out_channels + c_out) * out_h + row) * out_w + col;
output[out_idx] = acc;
}
} // namespace
// --- C++ Launcher ---
void conv2d_forward_wrapper(
torch::Tensor input,
torch::Tensor kernel,
torch::Tensor output) {
const int batch = static_cast<int>(input.size(0));
const int in_channels = static_cast<int>(input.size(1));
const int height = static_cast<int>(input.size(2));
const int width = static_cast<int>(input.size(3));
const int out_channels = static_cast<int>(kernel.size(0));
const int kernel_size = static_cast<int>(kernel.size(2));
const int out_h = height - kernel_size + 1;
const int out_w = width - kernel_size + 1;
// --- 专门化路径 (Fast Path) ---
// K4 C8 OC8
if (kernel_size == 4 && in_channels == 8 && out_channels == 8) {
const int channel_chunks = (out_channels + 8 - 1) / 8;
const dim3 grid = make_grid<12, 8>(out_w, out_h, batch, channel_chunks);
const dim3 block(12, 8);
conv2d_multi_static<4, 8, 8, 12, 8><<<grid, block>>>(
input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),
batch, out_channels, height, width);
return;
}
// K6 C16 OC16
if (kernel_size == 6 && in_channels == 16 && out_channels == 16) {
const int channel_chunks = (out_channels + 4 - 1) / 4;
const dim3 grid = make_grid<16, 10>(out_w, out_h, batch, channel_chunks);
const dim3 block(16, 10);
conv2d_multi_static<6, 16, 4, 16, 10><<<grid, block>>>(
input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),
batch, out_channels, height, width);
return;
}
// K8 C8 OC8
if (kernel_size == 8 && in_channels == 8 && out_channels == 8) {
const int channel_chunks = (out_channels + 4 - 1) / 4;
const dim3 grid = make_grid<16, 10>(out_w, out_h, batch, channel_chunks);
const dim3 block(16, 10);
conv2d_multi_static<8, 8, 4, 16, 10><<<grid, block>>>(
input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),
batch, out_channels, height, width);
return;
}
// K8 C16 OC16
if (kernel_size == 8 && in_channels == 16 && out_channels == 16) {
const int channel_chunks = (out_channels + 2 - 1) / 2;
const dim3 grid = make_grid<16, 8>(out_w, out_h, batch, channel_chunks);
const dim3 block(16, 8);
conv2d_multi_static<8, 16, 2, 16, 8><<<grid, block>>>(
input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),
batch, out_channels, height, width);
return;
}
// --- 通用 Fallback (Safe Path) ---
// 使用 Global Memory 版本,不申请 Shared Memory,避免 crash
dim3 block(16, 16);
dim3 grid(
(out_w + block.x - 1) / block.x,
(out_h + block.y - 1) / block.y,
batch * out_channels // Z 维度处理 Batch 和 OC
);
// 移除 smem_bytes 参数
conv2d_generic_kernel<<<grid, block>>>(
input.data_ptr<float>(),
kernel.data_ptr<float>(),
output.data_ptr<float>(),
batch,
in_channels,
out_channels,
height,
width,
kernel_size);
}
"""
cpp_source = (
"void conv2d_forward_wrapper("
" torch::Tensor input,"
" torch::Tensor kernel,"
" torch::Tensor output"
");"
)
conv2d_ext = load_inline(
name="conv2d_v5_robust",
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",
"-U__CUDA_NO_HALF_OPERATORS__", "-U__CUDA_NO_HALF_CONVERSIONS__"
],
verbose=False,
)
def custom_kernel(input_tuple):
input_tensor, kernel, output = input_tuple
if not input_tensor.is_contiguous(): input_tensor = input_tensor.contiguous()
if not kernel.is_contiguous(): kernel = kernel.contiguous()
if not output.is_contiguous(): output = output.contiguous()
conv2d_ext.conv2d_forward_wrapper(input_tensor, kernel, output)
return outputscrolls · 308 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 113167.
import torchfrom torch.utils.cpp_extension import load_inline- # --- 1. CUDA Kernel Source (核心逻辑不变,依然很快) ---+ # --- 1. CUDA Kernel Source ---cuda_source = r"""#include <torch/extension.h>#include <cuda_runtime.h>- #include <ATen/cuda/Exceptions.h>- #include <type_traits>+ #include <algorithm>- // 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.+ // 辅助函数:float4 累加+ template <int OC>+ __device__ __forceinline__ void accumulate_oc(const float* wptr, float v, float* acc) {+ if constexpr (OC % 4 == 0) {+ #pragma unroll+ for (int oc = 0; oc < OC; oc += 4) {+ const float4 w = reinterpret_cast<const float4*>(wptr)[oc / 4];+ acc[oc + 0] = fmaf(v, w.x, acc[oc + 0]);+ acc[oc + 1] = fmaf(v, w.y, acc[oc + 1]);+ acc[oc + 2] = fmaf(v, w.z, acc[oc + 2]);+ acc[oc + 3] = fmaf(v, w.w, acc[oc + 3]);+ }+ } else {+ #pragma unroll+ for (int oc = 0; oc < OC; ++oc) {+ acc[oc] = fmaf(v, wptr[oc], acc[oc]);+ }+ }+ }namespace {- constexpr int kMaxSharedBytes = 48 * 1024;+ template <int TW, int TH>+ inline dim3 make_grid(int out_w, int out_h, int batch, int channel_chunk) {+ const int grid_x = (out_w + TW - 1) / TW;+ const int grid_y = (out_h + TH - 1) / TH;+ const int grid_z = batch * channel_chunk;+ return dim3(grid_x, grid_y, grid_z);+ }- template <int K, int TileX, int TileY, int OC_PER_BLOCK>- __global__ void conv2d_grouped_kernel(+ // ---- 专门化多通道静态共享内存核 (保持不变,这部分逻辑是好的) ----+ template <int K, int C, int OC, int TW, int TH>+ __launch_bounds__(TW * TH, 2)+ __global__ void conv2d_multi_static(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;++ // ... 保持你的原始逻辑 ...+ // 为了节省篇幅,这里省略重复代码,但编译时需要包含你的原始 conv2d_multi_static 实现+ // 请确保这里包含你原代码中该函数的完整内容++ static_assert(TW * TH <= 1024, "block too large");+ constexpr int PATCH_W = TW + K - 1;+ constexpr int PATCH_H = TH + K - 1;+ constexpr int PATCH_STRIDE = (PATCH_W >= 15) ? PATCH_W : (PATCH_W + 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;+ const int out_w = width - K + 1;+ const int out_h = height - K + 1;- // 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 tile_x = blockIdx.x * TW;+ const int tile_y = blockIdx.y * TH;+ if (tile_x >= out_w || tile_y >= out_h) return;- if (b_idx >= batch) {- return;- }+ const int channel_chunks = (out_channels + OC - 1) / OC;+ const int bc = static_cast<int>(blockIdx.z);+ const int b_idx = bc / channel_chunks;+ const int chunk_idx = bc - b_idx * channel_chunks;+ const int c_start = chunk_idx * OC;+ if (b_idx >= batch || c_start >= out_channels) 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]+ alignas(16) __shared__ float k_smem[C][K][K][OC];+ alignas(16) __shared__ float in_smem[C][PATCH_H][PATCH_STRIDE];- const int tid = (threadIdx.z * blockDim.y + threadIdx.y) * blockDim.x + threadIdx.x;- const int threads = blockDim.x * blockDim.y * blockDim.z;+ const int tid = threadIdx.y * blockDim.x + threadIdx.x;+ const int blk_threads = blockDim.x * blockDim.y;- 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;+ const int total_k = OC * C * K * K;+ for (int idx = tid; idx < total_k; idx += blk_threads) {+ const int tmp = idx / OC;+ const int oc = idx - tmp * OC;+ const int ci = tmp / (K * K);+ const int rem0 = tmp - ci * K * K;+ const int kh = rem0 / K;+ const int kw = rem0 - kh * K;+ const int g_oc = c_start + oc;+ if (g_oc < out_channels) {+ const int k_idx = ((g_oc * C + ci) * K + kh) * K + kw;+ k_smem[ci][kh][kw][oc] = weight[k_idx];+ } else {+ k_smem[ci][kh][kw][oc] = 0.0f;}+ }- // 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;+ const int patch_elems = C * PATCH_H * PATCH_W;+ for (int idx = tid; idx < patch_elems; idx += blk_threads) {+ const int ci = idx / (PATCH_H * PATCH_W);+ const int rem0 = idx - ci * PATCH_H * PATCH_W;+ const int py = rem0 / PATCH_W;+ const int px = rem0 - py * PATCH_W;+ const int gx = tile_x + px;+ const int gy = tile_y + py;+ float val = 0.f;+ if (gx < width && gy < height && b_idx < batch) {+ const int in_idx = ((b_idx * C + ci) * height + gy) * width + gx;+ val = input[in_idx];}+ in_smem[ci][py][px] = val;+ }+ __syncthreads();- __syncthreads();+ const int out_x = tile_x + threadIdx.x;+ const int out_y = tile_y + threadIdx.y;+ if (out_x < out_w && out_y < out_h) {+ float acc[OC];+ #pragma unroll+ for (int oc = 0; oc < OC; ++oc) acc[oc] = 0.f;- 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;+ for (int ci = 0; ci < C; ++ci) {+ const float* in_ptr = &in_smem[ci][threadIdx.y][threadIdx.x];#pragma unroll- for (int kh = 0; kh < K; ++kh) {+ 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];- }- }+ for (int kw = 0; kw < K; ++kw) {+ const float v = in_ptr[kh * PATCH_STRIDE + kw];+ const float* wptr = &k_smem[ci][kh][kw][0];+ accumulate_oc<OC>(wptr, v, acc);}- const int out_idx = ((b_idx * out_channels + c_out) * out_h + h_out) * out_w + w_out;- output[out_idx] = acc;+ }}++ #pragma unroll+ for (int oc = 0; oc < OC; ++oc) {+ const int g_oc = c_start + oc;+ if (g_oc < out_channels) {+ const int out_idx = ((b_idx * out_channels + g_oc) * out_h + out_y) * out_w + out_x;+ output[out_idx] = acc[oc];+ }+ }+ }}- // Generic fallback when shared memory budget is exceeded or K is uncommon.- __global__ void conv2d_fallback_kernel(+ // ---- 【修复版】通用 fallback:Global Memory 版本 ----+ // 不使用 Shared Memory,避免因 channel 过多导致 Invalid Argument+ __global__ void conv2d_generic_kernel(const float* __restrict__ input,const float* __restrict__ weight,float* __restrict__ output,⋯ 3 unchanged linesint 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;++ // 每个 block 处理 16x16 的输出区域+ const int out_w = width - kernel_size + 1;+ const int out_h = height - kernel_size + 1;++ const int col = blockIdx.x * blockDim.x + threadIdx.x;+ const int row = blockIdx.y * blockDim.y + threadIdx.y;++ // Z 轴映射到 Batch 和 Out Channel+ const int bc = blockIdx.z;+ const int b = bc / out_channels;+ const int c_out = bc % out_channels; // 修正逻辑:bc = b * out_c + c_out- 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;- }+ if (col >= out_w || row >= out_h || b >= 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;- }+ float acc = 0.0f;- 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;+ // 直接遍历输入通道,利用 L1 Cache 进行缓存+ for (int c = 0; c < in_channels; ++c) {+ for (int kh = 0; kh < kernel_size; ++kh) {+ int in_h = row + kh;+ for (int kw = 0; kw < kernel_size; ++kw) {+ int in_w = col + kw;++ // 计算全局索引+ int in_idx = ((b * in_channels + c) * height + in_h) * width + in_w;+ int k_idx = ((c_out * in_channels + c) * kernel_size + kh) * kernel_size + kw;++ // 乘累加+ acc += input[in_idx] * weight[k_idx];+ }}+ }- 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;+ int out_idx = ((b * out_channels + c_out) * out_h + row) * out_w + col;+ output[out_idx] = acc;}- 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);+ // --- C++ Launcher ---+ void conv2d_forward_wrapper(+ torch::Tensor input,+ torch::Tensor kernel,+ torch::Tensor output) {+ const int batch = static_cast<int>(input.size(0));+ const int in_channels = static_cast<int>(input.size(1));+ const int height = static_cast<int>(input.size(2));+ const int width = static_cast<int>(input.size(3));- const int out_channels = kernel.size(0);- const int kernel_size = kernel.size(2);+ const int out_channels = static_cast<int>(kernel.size(0));+ const int kernel_size = static_cast<int>(kernel.size(2));- const int out_h = height - kernel_size + 1;- const int out_w = width - kernel_size + 1;+ 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);- };+ // --- 专门化路径 (Fast Path) ---+ // K4 C8 OC8+ if (kernel_size == 4 && in_channels == 8 && out_channels == 8) {+ const int channel_chunks = (out_channels + 8 - 1) / 8;+ const dim3 grid = make_grid<12, 8>(out_w, out_h, batch, channel_chunks);+ const dim3 block(12, 8);+ conv2d_multi_static<4, 8, 8, 12, 8><<<grid, block>>>(+ input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),+ batch, out_channels, height, width);+ return;+ }+ // K6 C16 OC16+ if (kernel_size == 6 && in_channels == 16 && out_channels == 16) {+ const int channel_chunks = (out_channels + 4 - 1) / 4;+ const dim3 grid = make_grid<16, 10>(out_w, out_h, batch, channel_chunks);+ const dim3 block(16, 10);+ conv2d_multi_static<6, 16, 4, 16, 10><<<grid, block>>>(+ input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),+ batch, out_channels, height, width);+ return;+ }+ // K8 C8 OC8+ if (kernel_size == 8 && in_channels == 8 && out_channels == 8) {+ const int channel_chunks = (out_channels + 4 - 1) / 4;+ const dim3 grid = make_grid<16, 10>(out_w, out_h, batch, channel_chunks);+ const dim3 block(16, 10);+ conv2d_multi_static<8, 8, 4, 16, 10><<<grid, block>>>(+ input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),+ batch, out_channels, height, width);+ return;+ }+ // K8 C16 OC16+ if (kernel_size == 8 && in_channels == 16 && out_channels == 16) {+ const int channel_chunks = (out_channels + 2 - 1) / 2;+ const dim3 grid = make_grid<16, 8>(out_w, out_h, batch, channel_chunks);+ const dim3 block(16, 8);+ conv2d_multi_static<8, 16, 2, 16, 8><<<grid, block>>>(+ input.data_ptr<float>(), kernel.data_ptr<float>(), output.data_ptr<float>(),+ batch, out_channels, height, width);+ return;+ }- 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();+ // --- 通用 Fallback (Safe Path) ---+ // 使用 Global Memory 版本,不申请 Shared Memory,避免 crash+ dim3 block(16, 16);+ dim3 grid(+ (out_w + block.x - 1) / block.x,+ (out_h + block.y - 1) / block.y,+ batch * out_channels // Z 维度处理 Batch 和 OC+ );++ // 移除 smem_bytes 参数+ conv2d_generic_kernel<<<grid, block>>>(+ input.data_ptr<float>(),+ kernel.data_ptr<float>(),+ output.data_ptr<float>(),+ batch,+ in_channels,+ out_channels,+ height,+ width,+ kernel_size);}"""⋯ 1 unchanged lines"void conv2d_forward_wrapper("" torch::Tensor input,"" torch::Tensor kernel,"- " torch::Tensor output" # 对应修改签名+ " torch::Tensor output"");")- # --- 2. 编译扩展 ---conv2d_ext = load_inline(- name="conv2d_v4_inplace",+ name="conv2d_v5_robust",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",+ "-O3", "--use_fast_math", "-std=c++17",+ "-U__CUDA_NO_HALF_OPERATORS__", "-U__CUDA_NO_HALF_CONVERSIONS__"],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 个 Tensorinput_tensor, kernel, output = input_tuple+ if not input_tensor.is_contiguous(): input_tensor = input_tensor.contiguous()+ if not kernel.is_contiguous(): kernel = kernel.contiguous()+ if not output.is_contiguous(): output = output.contiguous()- # 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 outputNo newline at end of file
scrolls · 596 diff lines total
Best evidence level for this revision: reported
JSON