Skip to content
KernelIndex
Search⌘K

submission 782583

ooousay · python · License unknown

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
2D convolutionsuite of 5 cases
NVIDIA A100
3.82ms
#1 of 40
2026-05-15

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 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
- 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