Skip to content
KernelIndex
Search⌘K

submission 526392

krasnaya_66854 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 134 lines, June 9 Researcher Reciprocity License v1.0.

test_ols3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-526392?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
5.04ms
#3 of 40
2026-03-10

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8cf6f026f7146606baaf621c61bab39e106ed6323765b1eaab5268ceb24d6827
license declaredunknown
license concludedunknown
authorskrasnaya_66854
imported2026-08-15

Kernel source

test_ols3.py134 lines
"""Overlap-save with tuned block size and no empty_cache overhead."""
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]
    # No empty_cache — avoids GPU stall
    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

    # Try bH=64 (2^6): L_h=33, 7x7=49 blocks. Smaller W_fft = 277MB → ~2.5ms cold.
    bH = _next_fast_len(kH + kH - 1)  # = next_fast_len(63) = 64
    bW = _next_fast_len(kW + kW - 1)  # = 64
    L_h = bH - kH + 1   # = 33
    L_w = bW - kW + 1   # = 33
    bW_half = bW // 2 + 1
    bhw = bH * bW_half   # = 64*33 = 2112

    n_h = (H_out + L_h - 1) // L_h  # = ceil(225/33) = 7
    n_w = (W_out + L_w - 1) // L_w  # = 7
    n_blocks = n_h * n_w  # = 49

    W_hw = _get_wfft(kernel, bH, bW)  # (bhw, Co, Ci) -- 277MB

    # Pad input at bottom/right for last blocks
    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
    x_blocks = x_padded.unfold(2, bH, L_h).unfold(3, bW, L_w)
    # (B, Ci, n_h, n_w, bH, bW) -> (n_h*n_w*B, Ci, bH, bW)
    x_blocks = x_blocks.permute(2, 3, 0, 1, 4, 5).contiguous().reshape(n_blocks * B, Ci, bH, bW)

    X_fft = torch.fft.rfft2(x_blocks, s=(bH, bW))  # (n_blocks*B, Ci, bH, bW_half)

    n_total = n_blocks * B
    X_hw = X_fft.reshape(n_total, Ci, bhw).permute(2, 1, 0).contiguous()  # (bhw, Ci, n_total)

    out_hw = torch.bmm(W_hw, X_hw)  # (bhw, Co, n_total)

    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)

    out_valid = out_blocks[:, :, :L_h, :L_w].contiguous()
    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 · 134 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 526371.

- """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)
- """
+ """Overlap-save with tuned block size and no empty_cache overhead."""
import os
os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'
⋯ 35 unchanged lines
while len(_cache_order) >= _MAX_ENTRIES:
oldest = _cache_order.pop(0)
del _wfft_cache[oldest]
- torch.cuda.empty_cache()
+ # No empty_cache — avoids GPU stall
Co, Ci, kH, kW = kernel.shape
bW_half = bW // 2 + 1
bhw = bH * bW_half
⋯ 9 unchanged lines
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
+ # Try bH=64 (2^6): L_h=33, 7x7=49 blocks. Smaller W_fft = 277MB → ~2.5ms cold.
+ bH = _next_fast_len(kH + kH - 1) # = next_fast_len(63) = 64
+ bW = _next_fast_len(kW + kW - 1) # = 64
+ L_h = bH - kH + 1 # = 33
+ L_w = bW - kW + 1 # = 33
bW_half = bW // 2 + 1
- bhw = bH * bW_half
+ bhw = bH * bW_half # = 64*33 = 2112
- n_h = (H_out + L_h - 1) // L_h
- n_w = (W_out + L_w - 1) // L_w
+ n_h = (H_out + L_h - 1) // L_h # = ceil(225/33) = 7
+ n_w = (W_out + L_w - 1) // L_w # = 7
+ n_blocks = n_h * n_w # = 49
- # Get kernel FFT (uses small bH×bW instead of large fH×fW)
- W_hw = _get_wfft(kernel, bH, bW) # (bhw, Co, Ci)
+ W_hw = _get_wfft(kernel, bH, bW) # (bhw, Co, Ci) -- 277MB
- # Pad input at BOTTOM/RIGHT so last blocks don't go out of bounds
- # Need: (n_h-1)*L_h + bH rows total
+ # Pad input at bottom/right for last blocks
total_h = (n_h - 1) * L_h + bH
total_w = (n_w - 1) * L_w + bW
pad_bot = max(0, total_h - H)
⋯ 3 unchanged lines
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)
+ # Extract all blocks via unfold
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)
+ # (B, Ci, n_h, n_w, bH, bW) -> (n_h*n_w*B, Ci, bH, bW)
+ x_blocks = x_blocks.permute(2, 3, 0, 1, 4, 5).contiguous().reshape(n_blocks * 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)
+ X_fft = torch.fft.rfft2(x_blocks, s=(bH, bW)) # (n_blocks*B, Ci, bH, bW_half)
- # Rearrange for bmm
- n_total = n_h * n_w * B
+ n_total = n_blocks * 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)
+ out_hw = torch.bmm(W_hw, X_hw) # (bhw, Co, n_total)
- # 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_valid = out_blocks[:, :, :L_h, :L_w].contiguous()
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)
scrolls · 97 diff lines total

Best evidence level for this revision: reported

JSON