submission 526371
krasnaya_66854 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 154 lines, June 9 Researcher Reciprocity License v1.0.
test_ols2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-526371?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:4024f08b7cbe3a99e26f52771dde2999b4bb67394a9f62775d15fa2eb6247392
license declaredunknown
license concludedunknown
authorskrasnaya_66854
imported2026-08-15
Kernel source
test_ols2.py154 lines
"""Overlap-save FFT cross-correlation (corrected).
For PyTorch conv2d (cross-correlation):
output[h,w] = sum_{dh,dw} input[h+dh, w+dw] * kernel[dh,dw]
Overlap-save for cross-correlation:
- Pad input at BOTTOM/RIGHT (not top/left)
- Blocks start at stride L = bH - kH + 1
- Valid output: FIRST L rows/cols of each block (not last)
"""
import os
os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'
import torch
import torch.nn.functional as F
from task import input_t, output_t
def _next_fast_len(n: int) -> int:
while True:
m = n
for p in (2, 3, 5):
while m % p == 0:
m //= p
if m == 1:
return n
n += 1
def _kernel_fingerprint(kernel: torch.Tensor):
n = kernel.numel()
f = kernel.ravel()
return (f[0].item(), f[n // 4].item(), f[n // 2].item(),
f[3 * n // 4].item(), f[-1].item())
_wfft_cache: dict = {}
_cache_order: list = []
_MAX_ENTRIES = 3
def _get_wfft(kernel: torch.Tensor, bH: int, bW: int) -> torch.Tensor:
fp = _kernel_fingerprint(kernel)
key = (fp, bH, bW)
if key in _wfft_cache:
_cache_order.remove(key)
_cache_order.append(key)
return _wfft_cache[key]
while len(_cache_order) >= _MAX_ENTRIES:
oldest = _cache_order.pop(0)
del _wfft_cache[oldest]
torch.cuda.empty_cache()
Co, Ci, kH, kW = kernel.shape
bW_half = bW // 2 + 1
bhw = bH * bW_half
W_fft_raw = torch.fft.rfft2(kernel.reshape(Co * Ci, kH, kW), s=(bH, bW))
W_fft = W_fft_raw.view(Co, Ci, bH, bW_half).conj().permute(2, 3, 0, 1).reshape(bhw, Co, Ci).contiguous()
_wfft_cache[key] = W_fft
_cache_order.append(key)
return W_fft
def _fft_conv2d_block(input_tensor, kernel, output):
B, Ci, H, W = input_tensor.shape
Co, _, kH, kW = kernel.shape
H_out, W_out = H - kH + 1, W - kW + 1
# Choose block size targeting ~4 blocks per dim
bH = _next_fast_len(max(H_out // 4 + kH - 1, kH))
bW = _next_fast_len(max(W_out // 4 + kW - 1, kW))
L_h = bH - kH + 1 # valid output rows per block
L_w = bW - kW + 1 # valid output cols per block
bW_half = bW // 2 + 1
bhw = bH * bW_half
n_h = (H_out + L_h - 1) // L_h
n_w = (W_out + L_w - 1) // L_w
# Get kernel FFT (uses small bH×bW instead of large fH×fW)
W_hw = _get_wfft(kernel, bH, bW) # (bhw, Co, Ci)
# Pad input at BOTTOM/RIGHT so last blocks don't go out of bounds
# Need: (n_h-1)*L_h + bH rows total
total_h = (n_h - 1) * L_h + bH
total_w = (n_w - 1) * L_w + bW
pad_bot = max(0, total_h - H)
pad_right = max(0, total_w - W)
if pad_bot > 0 or pad_right > 0:
x_padded = F.pad(input_tensor, (0, pad_right, 0, pad_bot))
else:
x_padded = input_tensor
# Extract all blocks via unfold (zero-copy view, stride=L_h/L_w)
# x_padded: (B, Ci, H+pad_bot, W+pad_right)
x_blocks = x_padded.unfold(2, bH, L_h).unfold(3, bW, L_w)
# x_blocks: (B, Ci, n_h, n_w, bH, bW)
x_blocks = x_blocks.permute(2, 3, 0, 1, 4, 5).contiguous()
# x_blocks: (n_h, n_w, B, Ci, bH, bW)
x_blocks = x_blocks.reshape(n_h * n_w * B, Ci, bH, bW)
# FFT of all blocks at once
X_fft = torch.fft.rfft2(x_blocks, s=(bH, bW)) # (n_h*n_w*B, Ci, bH, bW_half)
# Rearrange for bmm
n_total = n_h * n_w * B
X_hw = X_fft.reshape(n_total, Ci, bhw).permute(2, 1, 0).contiguous() # (bhw, Ci, n_total)
# bmm: (bhw, Co, Ci) x (bhw, Ci, n_total) -> (bhw, Co, n_total)
out_hw = torch.bmm(W_hw, X_hw)
# irfft2
out_fft = out_hw.permute(2, 1, 0).reshape(n_total, Co, bH, bW_half)
out_blocks = torch.fft.irfft2(out_fft, s=(bH, bW)) # (n_total, Co, bH, bW)
# Valid output: FIRST L_h x L_w of each block (cross-correlation valid region)
out_valid = out_blocks[:, :, :L_h, :L_w].contiguous() # (n_h*n_w*B, Co, L_h, L_w)
# Assemble: reshape to grid then to (B, Co, n_h*L_h, n_w*L_w)
out_grid = out_valid.reshape(n_h, n_w, B, Co, L_h, L_w)
out_grid = out_grid.permute(2, 3, 0, 4, 1, 5).contiguous()
out_full = out_grid.reshape(B, Co, n_h * L_h, n_w * L_w)
output[...] = out_full[:, :, :H_out, :W_out]
return output
def _fft_conv2d_full(input_tensor, kernel, output):
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)
fW_half = fW // 2 + 1
hw = fH * fW_half
X_fft = torch.fft.rfft2(input_tensor, s=(fH, fW))
W_hw = _get_wfft(kernel, fH, fW)
X_hw = X_fft.reshape(B, Ci, hw).permute(2, 1, 0).contiguous()
out_hw = torch.bmm(W_hw, X_hw)
out_fft = out_hw.permute(2, 1, 0).reshape(B, Co, fH, fW_half)
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 Ci <= 64:
return _fft_conv2d_full(input_tensor, kernel, output)
elif kW <= 16:
with torch.backends.cudnn.flags(deterministic=True, benchmark=True, allow_tf32=False):
output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
return output
else:
return _fft_conv2d_block(input_tensor, kernel, output)
scrolls · 154 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 526078.
+ """Overlap-save FFT cross-correlation (corrected).++ For PyTorch conv2d (cross-correlation):+ output[h,w] = sum_{dh,dw} input[h+dh, w+dw] * kernel[dh,dw]++ Overlap-save for cross-correlation:+ - Pad input at BOTTOM/RIGHT (not top/left)+ - Blocks start at stride L = bH - kH + 1+ - Valid output: FIRST L rows/cols of each block (not last)+ """import osos.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'⋯ 25 unchanged lines_MAX_ENTRIES = 3- def _get_wfft(kernel: torch.Tensor, fH: int, fW: int) -> torch.Tensor:- """Return W_fft in (hw, Co, Ci) layout with conj pre-applied.-- Stored as (fH*(fW//2+1), Co, Ci) complex64 — contiguous.- This layout allows efficient torch.bmm for the frequency-domain multiply.- """+ def _get_wfft(kernel: torch.Tensor, bH: int, bW: int) -> torch.Tensor:fp = _kernel_fingerprint(kernel)- key = (fp, fH, fW)+ key = (fp, bH, bW)if key in _wfft_cache:_cache_order.remove(key)_cache_order.append(key)⋯ 3 unchanged linesdel _wfft_cache[oldest]torch.cuda.empty_cache()Co, Ci, kH, kW = kernel.shape- fW_half = fW // 2 + 1- hw = fH * fW_half- # Compute W_fft and transpose to (hw, Co, Ci) layout with conj- W_fft_raw = torch.fft.rfft2(- kernel.reshape(Co * Ci, kH, kW), s=(fH, fW)- ) # (Co*Ci, fH, fW_half)- # Rearrange: (Co*Ci, fH, fW_half) → (Co, Ci, fH, fW_half) → conj → (hw, Co, Ci)- W_fft = W_fft_raw.view(Co, Ci, fH, fW_half).conj().permute(2, 3, 0, 1).reshape(hw, Co, Ci).contiguous()+ bW_half = bW // 2 + 1+ bhw = bH * bW_half+ W_fft_raw = torch.fft.rfft2(kernel.reshape(Co * Ci, kH, kW), s=(bH, bW))+ W_fft = W_fft_raw.view(Co, Ci, bH, bW_half).conj().permute(2, 3, 0, 1).reshape(bhw, Co, Ci).contiguous()_wfft_cache[key] = W_fft_cache_order.append(key)return W_fft- def _fft_conv2d(input_tensor, kernel, output):+ def _fft_conv2d_block(input_tensor, kernel, output):B, Ci, H, W = input_tensor.shapeCo, _, kH, kW = kernel.shapeH_out, W_out = H - kH + 1, W - kW + 1- fH = _next_fast_len(H + kH - 1)- fW = _next_fast_len(W + kW - 1)- fW_half = fW // 2 + 1- hw = fH * fW_half- X_fft = torch.fft.rfft2(input_tensor, s=(fH, fW)) # (B, Ci, fH, fW_half)- W_hw = _get_wfft(kernel, fH, fW) # (hw, Co, Ci)+ # Choose block size targeting ~4 blocks per dim+ bH = _next_fast_len(max(H_out // 4 + kH - 1, kH))+ bW = _next_fast_len(max(W_out // 4 + kW - 1, kW))+ L_h = bH - kH + 1 # valid output rows per block+ L_w = bW - kW + 1 # valid output cols per block+ bW_half = bW // 2 + 1+ bhw = bH * bW_half- # Rearrange X for bmm: (B, Ci, fH, fW_half) → (hw, Ci, B)- # Step 1: (B, Ci, fH, fW_half) → (B, Ci, hw) via reshape (Ci dim is contiguous after B)- # Step 2: permute to (hw, Ci, B) — need contiguous for bmm- X_hw = X_fft.reshape(B, Ci, hw).permute(2, 1, 0).contiguous() # (hw, Ci, B)+ n_h = (H_out + L_h - 1) // L_h+ n_w = (W_out + L_w - 1) // L_w- # bmm: (hw, Co, Ci) @ (hw, Ci, B) → (hw, Co, B)- out_hw = torch.bmm(W_hw, X_hw) # (hw, Co, B)+ # Get kernel FFT (uses small bH×bW instead of large fH×fW)+ W_hw = _get_wfft(kernel, bH, bW) # (bhw, Co, Ci)- # Rearrange back: (hw, Co, B) → (B, Co, fH, fW_half)- out_fft = out_hw.permute(2, 1, 0).reshape(B, Co, fH, fW_half)+ # Pad input at BOTTOM/RIGHT so last blocks don't go out of bounds+ # Need: (n_h-1)*L_h + bH rows total+ total_h = (n_h - 1) * L_h + bH+ total_w = (n_w - 1) * L_w + bW+ pad_bot = max(0, total_h - H)+ pad_right = max(0, total_w - W)+ if pad_bot > 0 or pad_right > 0:+ x_padded = F.pad(input_tensor, (0, pad_right, 0, pad_bot))+ else:+ x_padded = input_tensor+ # Extract all blocks via unfold (zero-copy view, stride=L_h/L_w)+ # x_padded: (B, Ci, H+pad_bot, W+pad_right)+ x_blocks = x_padded.unfold(2, bH, L_h).unfold(3, bW, L_w)+ # x_blocks: (B, Ci, n_h, n_w, bH, bW)+ x_blocks = x_blocks.permute(2, 3, 0, 1, 4, 5).contiguous()+ # x_blocks: (n_h, n_w, B, Ci, bH, bW)+ x_blocks = x_blocks.reshape(n_h * n_w * B, Ci, bH, bW)++ # FFT of all blocks at once+ X_fft = torch.fft.rfft2(x_blocks, s=(bH, bW)) # (n_h*n_w*B, Ci, bH, bW_half)++ # Rearrange for bmm+ n_total = n_h * n_w * B+ X_hw = X_fft.reshape(n_total, Ci, bhw).permute(2, 1, 0).contiguous() # (bhw, Ci, n_total)++ # bmm: (bhw, Co, Ci) x (bhw, Ci, n_total) -> (bhw, Co, n_total)+ out_hw = torch.bmm(W_hw, X_hw)++ # irfft2+ out_fft = out_hw.permute(2, 1, 0).reshape(n_total, Co, bH, bW_half)+ out_blocks = torch.fft.irfft2(out_fft, s=(bH, bW)) # (n_total, Co, bH, bW)++ # Valid output: FIRST L_h x L_w of each block (cross-correlation valid region)+ out_valid = out_blocks[:, :, :L_h, :L_w].contiguous() # (n_h*n_w*B, Co, L_h, L_w)++ # Assemble: reshape to grid then to (B, Co, n_h*L_h, n_w*L_w)+ out_grid = out_valid.reshape(n_h, n_w, B, Co, L_h, L_w)+ out_grid = out_grid.permute(2, 3, 0, 4, 1, 5).contiguous()+ out_full = out_grid.reshape(B, Co, n_h * L_h, n_w * L_w)+ output[...] = out_full[:, :, :H_out, :W_out]+ return output+++ def _fft_conv2d_full(input_tensor, kernel, output):+ 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)+ fW_half = fW // 2 + 1+ hw = fH * fW_half+ X_fft = torch.fft.rfft2(input_tensor, s=(fH, fW))+ W_hw = _get_wfft(kernel, fH, fW)+ X_hw = X_fft.reshape(B, Ci, hw).permute(2, 1, 0).contiguous()+ out_hw = torch.bmm(W_hw, X_hw)+ out_fft = out_hw.permute(2, 1, 0).reshape(B, Co, fH, fW_half)output[...] = torch.fft.irfft2(out_fft, s=(fH, fW))[:, :, :H_out, :W_out]return output⋯ 1 unchanged linesdef custom_kernel(data: input_t) -> output_t:input_tensor, kernel, output = dataCo, Ci, kH, kW = kernel.shape-if Ci <= 64:- return _fft_conv2d(input_tensor, kernel, output)+ return _fft_conv2d_full(input_tensor, kernel, output)elif kW <= 16:with torch.backends.cudnn.flags(deterministic=True, benchmark=True, allow_tf32=False):output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)return outputelse:- return _fft_conv2d(input_tensor, kernel, output)+ return _fft_conv2d_block(input_tensor, kernel, output)
scrolls · 160 diff lines total
Best evidence level for this revision: reported
JSON