submission 524332
krasnaya_66854 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 55 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-524332?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:14acb6a967a3ed2a92f80f447c5489901bb51595fe9a5a79f58766373406f078
license declaredunknown
license concludedunknown
authorskrasnaya_66854
imported2026-08-15
Kernel source
submission.py55 lines
import os
os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'
import torch
import torch.nn.functional as F
from task import input_t, output_t
from utils import DeterministicContext
def _next_fast_len(n: int) -> int:
"""Smallest integer >= n with only prime factors 2, 3, 5 (5-smooth = fast for cuFFT)."""
while True:
m = n
for p in (2, 3, 5):
while m % p == 0:
m //= p
if m == 1:
return n
n += 1
def _fft_conv2d(input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor) -> torch.Tensor:
B, Ci, H, W = input_tensor.shape
Co, _, kH, kW = kernel.shape
H_out, W_out = H - kH + 1, W - kW + 1
fH = _next_fast_len(H + kH - 1)
fW = _next_fast_len(W + kW - 1)
fW2 = fW // 2 + 1
X_fft = torch.fft.rfft2(input_tensor, s=(fH, fW))
W_fft = torch.fft.rfft2(
kernel.reshape(Co * Ci, kH, kW), s=(fH, fW)
).reshape(Co, Ci, fH, fW2)
out_fft = torch.einsum('bihw,oihw->bohw', X_fft, W_fft.conj())
output[...] = torch.fft.irfft2(out_fft, s=(fH, fW))[:, :, :H_out, :W_out]
return output
def custom_kernel(data: input_t) -> output_t:
input_tensor, kernel, output = data
Co, Ci, kH, kW = kernel.shape
if kW <= 16:
# FFT convolution. For k<=16, FFT error is ~2e-5 (well within 1e-3 tolerance).
# W_fft is at most (128^2) * 288 * 145 * 8 = 5.1 GB which fits in A100 80GB.
return _fft_conv2d(input_tensor, kernel, output)
else:
# k=32: FFT float32 error (~4e-3) exceeds atol=1e-3 for near-zero elements.
# Must use DeterministicContext to reproduce the reference's exact float32 result.
with DeterministicContext():
output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
return output
scrolls · 55 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 524244.
+ import os+ os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'+import torchimport torch.nn.functional as Ffrom task import input_t, output_t⋯ 12 unchanged linesn += 1- # FFT is fast when W_fft tensor is small. W_fft size = Co*Ci * fH * fW2 * 8 bytes.- # For C=128, k=16: 16384 * 288 * 145 * 8 = 5.1 GB -- memory bandwidth kills the gain.- # For C=64, k=16: 4096 * 144 * 73 * 8 = 344 MB -- fast.- # Threshold: only use FFT when Co*Ci*fH*fW2 * 8 bytes < ~512 MB.- _FFT_MAX_KSIZE = 16- _FFT_MAX_CHANNELS = 64 # above this, W_fft memory cost dominates--def _fft_conv2d(input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor) -> torch.Tensor:B, Ci, H, W = input_tensor.shapeCo, _, kH, kW = kernel.shapeH_out, W_out = H - kH + 1, W - kW + 1- # FFT size must be >= H+kH-1 to compute linear (non-circular) convolutionfH = _next_fast_len(H + kH - 1)fW = _next_fast_len(W + kW - 1)fW2 = fW // 2 + 1- # Input FFT: (B, Ci, H, W) -> (B, Ci, fH, fW2) complex64X_fft = torch.fft.rfft2(input_tensor, s=(fH, fW))-- # Kernel FFT: batch all Co*Ci kernels in one cuFFT call- # (Co, Ci, kH, kW) -> (Co*Ci, kH, kW) -> rfft2 -> (Co, Ci, fH, fW2) complex64W_fft = torch.fft.rfft2(kernel.reshape(Co * Ci, kH, kW), s=(fH, fW)).reshape(Co, Ci, fH, fW2)- # Cross-correlation in frequency domain:- # out[b, o, h, w] = sum_i X[b, i, h, w] * conj(W[o, i, h, w])out_fft = torch.einsum('bihw,oihw->bohw', X_fft, W_fft.conj())-- # IFFT and slice valid region- result = torch.fft.irfft2(out_fft, s=(fH, fW))- output[...] = result[:, :, :H_out, :W_out]+ output[...] = torch.fft.irfft2(out_fft, s=(fH, fW))[:, :, :H_out, :W_out]return output⋯ 1 unchanged linesinput_tensor, kernel, output = dataCo, Ci, kH, kW = kernel.shape- use_fft = (kW <= _FFT_MAX_KSIZE) and (Ci <= _FFT_MAX_CHANNELS)- if use_fft:- # FFT convolution: fast when W_fft fits in memory (~344 MB for C=64,k=16)+ if kW <= 16:+ # FFT convolution. For k<=16, FFT error is ~2e-5 (well within 1e-3 tolerance).+ # W_fft is at most (128^2) * 288 * 145 * 8 = 5.1 GB which fits in A100 80GB.return _fft_conv2d(input_tensor, kernel, output)- elif kW <= _FFT_MAX_KSIZE:- # Large channels (C=128, k=16): W_fft would be 5.1 GB.- # Plain cuDNN is faster here.- output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)- return outputelse:- # Large kernels (k=32): FFT float32 error too high; must match reference.+ # k=32: FFT float32 error (~4e-3) exceeds atol=1e-3 for near-zero elements.+ # Must use DeterministicContext to reproduce the reference's exact float32 result.with DeterministicContext():output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)return output
scrolls · 71 diff lines total
Best evidence level for this revision: reported
JSON