Skip to content
KernelIndex
Search⌘K

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
2D convolutionsuite of 5 cases
NVIDIA B200
306.3ms
#27 of 28
2025-11-30

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-memoryalignas(16) __shared__ float k_smem[C][K][K][OC];
vector-width = float4const 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 output
scrolls · 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 torch
from 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 lines
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;
+
+ // 每个 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 个 Tensor
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()
- # 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
No newline at end of file
scrolls · 596 diff lines total

Best evidence level for this revision: reported

JSON