Skip to content
KernelIndex
Search⌘K

submission 538107

krasnaya_66854 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_b200.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-538107?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
2D convolutionsuite of 5 cases
NVIDIA B200
42.3ms
#17 of 28
2026-03-12

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0f0eb25b92d882a78971dbd00e47699f34cf364c95df11b4e5390602ae7a2e09
license declaredunknown
license concludedunknown
authorskrasnaya_66854
imported2026-08-15

Kernel source

submission_b200.py203 lines
"""Shared conv2d implementations for the fixed conv2d_v2 benchmark set."""

import os

os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"

import torch
import torch.nn.functional as F

from task import input_t, output_t

ShapeKey = tuple[int, int, int, int]
PlanValue = tuple[str, int | tuple[int, int] | None]

_WFFT_CACHE: dict[tuple[tuple[float, ...], int, int], torch.Tensor] = {}
_CACHE_ORDER: list[tuple[tuple[float, ...], int, int]] = []
_MAX_ENTRIES = 4


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) -> tuple[float, ...]:
    flat = kernel.reshape(-1)
    n = flat.numel()
    idxs = (0, n // 7, n // 3, n // 2, (5 * n) // 7, n - 1)
    return tuple(float(flat[i].item()) for i in idxs)


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]

    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 conv2d_cudnn(
    input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor
) -> torch.Tensor:
    with torch.backends.cudnn.flags(
        deterministic=True,
        benchmark=True,
        allow_tf32=False,
    ):
        output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
    return output


def conv2d_cudnn_channels_last(
    input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor
) -> torch.Tensor:
    x = input_tensor.contiguous(memory_format=torch.channels_last)
    w = kernel.contiguous(memory_format=torch.channels_last)
    with torch.backends.cudnn.flags(
        deterministic=True,
        benchmark=True,
        allow_tf32=False,
    ):
        y = F.conv2d(x, w, stride=1, padding=0)
    output[...] = y.contiguous()
    return output


def fft_conv2d_full(
    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)
    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 fft_conv2d_block(
    input_tensor: torch.Tensor,
    kernel: torch.Tensor,
    output: torch.Tensor,
    bH: int,
    bW: int,
) -> torch.Tensor:
    B, Ci, H, W = input_tensor.shape
    Co, _, kH, kW = kernel.shape
    H_out, W_out = H - kH + 1, W - kW + 1

    bH = next_fast_len(max(bH, kH))
    bW = next_fast_len(max(bW, kW))
    L_h = bH - kH + 1
    L_w = bW - kW + 1
    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
    n_blocks = n_h * n_w

    W_hw = get_wfft(kernel, bH, bW)

    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)
    x_padded = (
        F.pad(input_tensor, (0, pad_right, 0, pad_bot))
        if pad_bot > 0 or pad_right > 0
        else input_tensor
    )

    x_blocks = x_padded.unfold(2, bH, L_h).unfold(3, bW, L_w)
    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_total = n_blocks * B
    X_hw = X_fft.reshape(n_total, Ci, bhw).permute(2, 1, 0).contiguous()
    out_hw = torch.bmm(W_hw, X_hw)

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

    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 shape_key(input_tensor: torch.Tensor, kernel: torch.Tensor) -> ShapeKey:
    B, C, H, _ = input_tensor.shape
    _, _, K, _ = kernel.shape
    return (B, C, H, K)


def run_plan(
    data: input_t, plan: dict[ShapeKey, PlanValue]
) -> output_t:
    input_tensor, kernel, output = data
    shape = shape_key(input_tensor, kernel)
    method, param = plan.get(shape, ("fft_block", (64, 64)))

    if method == "cudnn":
        return conv2d_cudnn(input_tensor, kernel, output)
    if method == "cudnn_cl":
        return conv2d_cudnn_channels_last(input_tensor, kernel, output)
    if method == "fft_full":
        return fft_conv2d_full(input_tensor, kernel, output)

    if isinstance(param, tuple):
        bH, bW = param
    else:
        bH = bW = param or 64
    return fft_conv2d_block(input_tensor, kernel, output, bH, bW)


PLAN = {(4, 64, 128, 8): ('cudnn_cl', None), (4, 64, 128, 16): ('cudnn_cl', None), (2, 128, 256, 16): ('fft_full', None), (1, 128, 256, 32): ('cudnn_cl', None)}


def custom_kernel(data: input_t) -> output_t:
    return run_plan(data, PLAN)
scrolls · 203 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 526573.

- """Overlap-save with tuned block size and no empty_cache overhead."""
+ """Shared conv2d implementations for the fixed conv2d_v2 benchmark set."""
+
import os
- os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'
+ os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
+
import torch
import torch.nn.functional as F
+
from task import input_t, output_t
+ ShapeKey = tuple[int, int, int, int]
+ PlanValue = tuple[str, int | tuple[int, int] | None]
- def _next_fast_len(n: int) -> int:
+ _WFFT_CACHE: dict[tuple[tuple[float, ...], int, int], torch.Tensor] = {}
+ _CACHE_ORDER: list[tuple[tuple[float, ...], int, int]] = []
+ _MAX_ENTRIES = 4
+
+
+ def next_fast_len(n: int) -> int:
while True:
m = n
for p in (2, 3, 5):
⋯ 4 unchanged lines
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())
+ def kernel_fingerprint(kernel: torch.Tensor) -> tuple[float, ...]:
+ flat = kernel.reshape(-1)
+ n = flat.numel()
+ idxs = (0, n // 7, n // 3, n // 2, (5 * n) // 7, n - 1)
+ return tuple(float(flat[i].item()) for i in idxs)
- _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]
- 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)
+ 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):
+ def conv2d_cudnn(
+ input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor
+ ) -> torch.Tensor:
+ with torch.backends.cudnn.flags(
+ deterministic=True,
+ benchmark=True,
+ allow_tf32=False,
+ ):
+ output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
+ return output
+
+
+ def conv2d_cudnn_channels_last(
+ input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor
+ ) -> torch.Tensor:
+ x = input_tensor.contiguous(memory_format=torch.channels_last)
+ w = kernel.contiguous(memory_format=torch.channels_last)
+ with torch.backends.cudnn.flags(
+ deterministic=True,
+ benchmark=True,
+ allow_tf32=False,
+ ):
+ y = F.conv2d(x, w, stride=1, padding=0)
+ output[...] = y.contiguous()
+ return output
+
+
+ def fft_conv2d_full(
+ 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
- # 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
+ 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 fft_conv2d_block(
+ input_tensor: torch.Tensor,
+ kernel: torch.Tensor,
+ output: torch.Tensor,
+ bH: int,
+ bW: int,
+ ) -> torch.Tensor:
+ B, Ci, H, W = input_tensor.shape
+ Co, _, kH, kW = kernel.shape
+ H_out, W_out = H - kH + 1, W - kW + 1
+
+ bH = next_fast_len(max(bH, kH))
+ bW = next_fast_len(max(bW, kW))
+ L_h = bH - kH + 1
+ L_w = bW - kW + 1
bW_half = bW // 2 + 1
- bhw = bH * bW_half # = 64*33 = 2112
+ bhw = bH * bW_half
- 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
+ n_h = (H_out + L_h - 1) // L_h
+ n_w = (W_out + L_w - 1) // L_w
+ n_blocks = n_h * n_w
- W_hw = _get_wfft(kernel, bH, bW) # (bhw, Co, Ci) -- 277MB
+ W_hw = get_wfft(kernel, bH, bW)
- # 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
+ x_padded = (
+ F.pad(input_tensor, (0, pad_right, 0, pad_bot))
+ if pad_bot > 0 or pad_right > 0
+ else 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_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)
-
+ X_fft = torch.fft.rfft2(x_blocks, s=(bH, bW))
n_total = n_blocks * B
- X_hw = X_fft.reshape(n_total, Ci, bhw).permute(2, 1, 0).contiguous() # (bhw, Ci, n_total)
+ X_hw = X_fft.reshape(n_total, Ci, bhw).permute(2, 1, 0).contiguous()
+ out_hw = torch.bmm(W_hw, X_hw)
- 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_blocks = torch.fft.irfft2(out_fft, s=(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)
⋯ 3 unchanged lines
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 shape_key(input_tensor: torch.Tensor, kernel: torch.Tensor) -> ShapeKey:
+ B, C, H, _ = input_tensor.shape
+ _, _, K, _ = kernel.shape
+ return (B, C, H, K)
- def custom_kernel(data: input_t) -> output_t:
+ def run_plan(
+ data: input_t, plan: dict[ShapeKey, PlanValue]
+ ) -> 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
+ shape = shape_key(input_tensor, kernel)
+ method, param = plan.get(shape, ("fft_block", (64, 64)))
+
+ if method == "cudnn":
+ return conv2d_cudnn(input_tensor, kernel, output)
+ if method == "cudnn_cl":
+ return conv2d_cudnn_channels_last(input_tensor, kernel, output)
+ if method == "fft_full":
+ return fft_conv2d_full(input_tensor, kernel, output)
+
+ if isinstance(param, tuple):
+ bH, bW = param
else:
- return _fft_conv2d_block(input_tensor, kernel, output)
+ bH = bW = param or 64
+ return fft_conv2d_block(input_tensor, kernel, output, bH, bW)
+
+
+ PLAN = {(4, 64, 128, 8): ('cudnn_cl', None), (4, 64, 128, 16): ('cudnn_cl', None), (2, 128, 256, 16): ('fft_full', None), (1, 128, 256, 32): ('cudnn_cl', None)}
+
+
+ def custom_kernel(data: input_t) -> output_t:
+ return run_plan(data, PLAN)
scrolls · 271 diff lines total

Best evidence level for this revision: reported

JSON