submission 782583
ooousay · python · License unknown
Kernel source · 117 lines ↓holds 1 record
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 117 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-782583?include=source"interfacepython
Compatibility
measured onNVIDIA A100
declared hardwareNVIDIA A100
architecturessm_80
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:1f2f90694afbecdc63a1ac6b071373b54e0f11e8ec5bc20bac150960bad1207e
license declaredunknown
license concludedunknown
authorsooousay
imported2026-08-15
Kernel source
submission.py117 lines
"""Tiled overlap-save FFT-conv via cuFFT + complex bmm.
The conv2d benchmark shapes have square kernels (K=8/16/32), no stride/dilation,
no padding, FP32, with rtol=atol=1e-3. Direct conv is compute-bound at these
sizes; cuDNN responds by picking an FFT-tiling path internally (engine 38/39/46
family, `ampere_gcgemm_64x32_tn`), but it pays for FFT pad/transform/transpose
as separate HBM-roundtripping kernels (~16 ms of 22 ms on shape 2).
This submission does the same tiled overlap-save FFT-conv but in PyTorch +
cuFFT + cuBLAS complex bmm, with per-shape FFT tile size T_in. The tiling
ensures input/weight spectra stay small (no full Co·Ci·H·(W/2+1) spectral
kernel materialisation — that was the trap `experiments/fft_proto.py` fell
into and got 16 GB scratch on shape 2).
Algorithm (cross-correlation overlap-save):
1. Pad input to (Nt_h·T_o + K-1, Nt_w·T_o + K-1) where T_o = T_in - K + 1.
2. View as (B, C, Nt_h, Nt_w, T_in, T_in) via as_strided (overlap K-1).
3. rfft2 along last two dims → (B, C, Nt_h, Nt_w, T_in, T_in//2+1) complex.
4. rfft2 of zero-padded weight → (Cout, Cin, T_in, T_in//2+1); take conj
because F.conv2d is cross-correlation, not convolution.
5. Per-frequency-bin complex bmm: (B·Nt_h·Nt_w, Cin) × (Cin, Cout) for each
of T_in·(T_in//2+1) bins. cuBLAS handles this in one batched call.
6. irfft2 → (B, Cout, Nt_h, Nt_w, T_in, T_in).
7. Crop the first T_o×T_o per tile (valid cross-correlation region) and
unspread back to (B, Cout, H_out, W_out).
Per-shape T_in was swept (`outputs/fft_tin_sweep_einsum.log`) on this A100:
- S128 K8 C64 B4: T_in=32 → 0.69 ms vs cuDNN 0.98 (1.41×)
- S128 K16 C64 B4: T_in=40 → 1.05 ms vs cuDNN 1.58 (1.50×)
- S256 K16 C128 B2: T_in=42 → 5.13 ms vs cuDNN 22.52 (4.39×)
- S256 K32 C128 B1: T_in=64 → 3.91 ms vs cuDNN 19.03 (4.86×)
Bench total: 15.91 ms vs cuDNN 66.63 ms (4.19×, 100-iter run).
Why T_in is not just K+1 (smallest valid):
- Too small → many tiles → bmm batch dim huge, small M/N per matmul → poor
tensor-core utilisation. Bigger tiles trade FFT cost for better GEMM
utilisation, with a per-shape sweet spot.
Correctness: passes rtol=atol=1e-3 vs cuDNN strict-FP32 reference on all 5
benchmark shapes across 50 leaderboard-style regenerated-input iters.
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from task import input_t, output_t
torch.backends.cudnn.allow_tf32 = False
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.benchmark = True
torch.backends.cudnn.deterministic = False
def _fft_conv2d_tiled(inp: torch.Tensor, wt: torch.Tensor, out: torch.Tensor,
T_in: int) -> torch.Tensor:
B, C, H, W = inp.shape
Cout, _, K, _ = wt.shape
T_o = T_in - K + 1
H_out = H - K + 1
W_out = W - K + 1
Nt_h = (H_out + T_o - 1) // T_o
Nt_w = (W_out + T_o - 1) // T_o
H_pad = Nt_h * T_o + K - 1
W_pad = Nt_w * T_o + K - 1
if H_pad > H or W_pad > W:
inp_p = F.pad(inp, (0, W_pad - W, 0, H_pad - H))
else:
inp_p = inp
s = inp_p.stride()
inp_tiles = inp_p.as_strided(
size=(B, C, Nt_h, Nt_w, T_in, T_in),
stride=(s[0], s[1], T_o * s[2], T_o * s[3], s[2], s[3]),
)
inp_freq = torch.fft.rfft2(inp_tiles, s=(T_in, T_in))
F_w = T_in // 2 + 1
wt_p = F.pad(wt, (0, T_in - K, 0, T_in - K))
wt_freq = torch.fft.rfft2(wt_p, s=(T_in, T_in)).conj()
# einsum contracts C; keeps (B, Nt_h, Nt_w, T_in, F_w) as batch and Cout as
# output. PyTorch's einsum planner picks a smarter contraction order than
# the explicit permute+bmm version, especially for shape 4 where it saves
# ~1.5 ms on the (T_in, F_w) batched complex GEMM.
out_freq = torch.einsum('bcijkl,ockl->boijkl', inp_freq, wt_freq)
out_tiles = torch.fft.irfft2(out_freq, s=(T_in, T_in))
valid = out_tiles[:, :, :, :, :T_o, :T_o]
valid = valid.permute(0, 1, 2, 4, 3, 5).contiguous().view(
B, Cout, Nt_h * T_o, Nt_w * T_o
)
out.copy_(valid[:, :, :H_out, :W_out])
return out
# Per-shape T_in tuned via the sweep in outputs/fft_tin_sweep.log. Keyed on
# (KSZ, C_in). Shapes not in the table fall through to cuDNN.
_T_IN_FOR_SHAPE: dict[tuple[int, int], int] = {
(8, 64): 32,
(16, 64): 40,
(16, 128): 42,
(32, 128): 64,
}
def custom_kernel(data: input_t) -> output_t:
input_tensor, kernel, output = data
_, C, _, _ = input_tensor.shape
KSZ = kernel.shape[-1]
T_in = _T_IN_FOR_SHAPE.get((KSZ, C))
if T_in is not None:
_fft_conv2d_tiled(input_tensor, kernel, output, T_in)
else:
output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
return output
scrolls · 117 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 782512.
- """Current best submission.+ """Tiled overlap-save FFT-conv via cuFFT + complex bmm.- cuDNN baseline + `cudnn.benchmark = True`.+ The conv2d benchmark shapes have square kernels (K=8/16/32), no stride/dilation,+ no padding, FP32, with rtol=atol=1e-3. Direct conv is compute-bound at these+ sizes; cuDNN responds by picking an FFT-tiling path internally (engine 38/39/46+ family, `ampere_gcgemm_64x32_tn`), but it pays for FFT pad/transform/transpose+ as separate HBM-roundtripping kernels (~16 ms of 22 ms on shape 2).- Root cause of the original 2.59 s outlier on S256 K32 C128 B1: with- `cudnn.deterministic = True` (set by the reference's DeterministicContext) and- `benchmark = False` (PyTorch default), cuDNN picks an FFT-conv engine that- degenerates to a memory-bound mat-vec at B=1.+ This submission does the same tiled overlap-save FFT-conv but in PyTorch ++ cuFFT + cuBLAS complex bmm, with per-shape FFT tile size T_in. The tiling+ ensures input/weight spectra stay small (no full Co·Ci·H·(W/2+1) spectral+ kernel materialisation — that was the trap `experiments/fft_proto.py` fell+ into and got 16 GB scratch on shape 2).- Flipping `benchmark = True` lets cuDNN run a timed heuristic on the first call- and select a much faster implicit-GEMM path for K=32. The eval harness does- one warmup call before timing, so algo selection happens off-clock.+ Algorithm (cross-correlation overlap-save):+ 1. Pad input to (Nt_h·T_o + K-1, Nt_w·T_o + K-1) where T_o = T_in - K + 1.+ 2. View as (B, C, Nt_h, Nt_w, T_in, T_in) via as_strided (overlap K-1).+ 3. rfft2 along last two dims → (B, C, Nt_h, Nt_w, T_in, T_in//2+1) complex.+ 4. rfft2 of zero-padded weight → (Cout, Cin, T_in, T_in//2+1); take conj+ because F.conv2d is cross-correlation, not convolution.+ 5. Per-frequency-bin complex bmm: (B·Nt_h·Nt_w, Cin) × (Cin, Cout) for each+ of T_in·(T_in//2+1) bins. cuBLAS handles this in one batched call.+ 6. irfft2 → (B, Cout, Nt_h, Nt_w, T_in, T_in).+ 7. Crop the first T_o×T_o per tile (valid cross-correlation region) and+ unspread back to (B, Cout, H_out, W_out).- Tradeoff: `benchmark = True` can pick worse algos for some mid-sized shapes- when L2 is cleared between calls (the heuristic times candidates with hot L2).- On this benchmark suite, that costs ~30 ms across all shapes vs cuDNN-default- but saves ~2470 ms on the K=32 shape — net ~11x speedup on total wall time.+ Per-shape T_in was swept (`outputs/fft_tin_sweep_einsum.log`) on this A100:+ - S128 K8 C64 B4: T_in=32 → 0.69 ms vs cuDNN 0.98 (1.41×)+ - S128 K16 C64 B4: T_in=40 → 1.05 ms vs cuDNN 1.58 (1.50×)+ - S256 K16 C128 B2: T_in=42 → 5.13 ms vs cuDNN 22.52 (4.39×)+ - S256 K32 C128 B1: T_in=64 → 3.91 ms vs cuDNN 19.03 (4.86×)+ Bench total: 15.91 ms vs cuDNN 66.63 ms (4.19×, 100-iter run).- TF32 stays off: reference is strict FP32, default-TF32 cuDNN diverges >1e-3- on K=16 and K=32 shapes.+ Why T_in is not just K+1 (smallest valid):+ - Too small → many tiles → bmm batch dim huge, small M/N per matmul → poor+ tensor-core utilisation. Bigger tiles trade FFT cost for better GEMM+ utilisation, with a per-shape sweet spot.++ Correctness: passes rtol=atol=1e-3 vs cuDNN strict-FP32 reference on all 5+ benchmark shapes across 50 leaderboard-style regenerated-input iters."""- from task import input_t, output_t+ from __future__ import annotationsimport torchimport torch.nn.functional as F+ from task import input_t, output_t+torch.backends.cudnn.allow_tf32 = Falsetorch.backends.cuda.matmul.allow_tf32 = Falsetorch.backends.cudnn.benchmark = Truetorch.backends.cudnn.deterministic = False+ def _fft_conv2d_tiled(inp: torch.Tensor, wt: torch.Tensor, out: torch.Tensor,+ T_in: int) -> torch.Tensor:+ B, C, H, W = inp.shape+ Cout, _, K, _ = wt.shape+ T_o = T_in - K + 1+ H_out = H - K + 1+ W_out = W - K + 1+ Nt_h = (H_out + T_o - 1) // T_o+ Nt_w = (W_out + T_o - 1) // T_o+ H_pad = Nt_h * T_o + K - 1+ W_pad = Nt_w * T_o + K - 1++ if H_pad > H or W_pad > W:+ inp_p = F.pad(inp, (0, W_pad - W, 0, H_pad - H))+ else:+ inp_p = inp++ s = inp_p.stride()+ inp_tiles = inp_p.as_strided(+ size=(B, C, Nt_h, Nt_w, T_in, T_in),+ stride=(s[0], s[1], T_o * s[2], T_o * s[3], s[2], s[3]),+ )+ inp_freq = torch.fft.rfft2(inp_tiles, s=(T_in, T_in))+ F_w = T_in // 2 + 1++ wt_p = F.pad(wt, (0, T_in - K, 0, T_in - K))+ wt_freq = torch.fft.rfft2(wt_p, s=(T_in, T_in)).conj()++ # einsum contracts C; keeps (B, Nt_h, Nt_w, T_in, F_w) as batch and Cout as+ # output. PyTorch's einsum planner picks a smarter contraction order than+ # the explicit permute+bmm version, especially for shape 4 where it saves+ # ~1.5 ms on the (T_in, F_w) batched complex GEMM.+ out_freq = torch.einsum('bcijkl,ockl->boijkl', inp_freq, wt_freq)++ out_tiles = torch.fft.irfft2(out_freq, s=(T_in, T_in))+ valid = out_tiles[:, :, :, :, :T_o, :T_o]+ valid = valid.permute(0, 1, 2, 4, 3, 5).contiguous().view(+ B, Cout, Nt_h * T_o, Nt_w * T_o+ )+ out.copy_(valid[:, :, :H_out, :W_out])+ return out+++ # Per-shape T_in tuned via the sweep in outputs/fft_tin_sweep.log. Keyed on+ # (KSZ, C_in). Shapes not in the table fall through to cuDNN.+ _T_IN_FOR_SHAPE: dict[tuple[int, int], int] = {+ (8, 64): 32,+ (16, 64): 40,+ (16, 128): 42,+ (32, 128): 64,+ }++def custom_kernel(data: input_t) -> output_t:input_tensor, kernel, output = data- output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)+ _, C, _, _ = input_tensor.shape+ KSZ = kernel.shape[-1]+ T_in = _T_IN_FOR_SHAPE.get((KSZ, C))+ if T_in is not None:+ _fft_conv2d_tiled(input_tensor, kernel, output, T_in)+ else:+ output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)return output
scrolls · 133 diff lines total
Best evidence level for this revision: reported
JSON