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
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 torchimport 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 = nfor p in (2, 3, 5):⋯ 4 unchanged linesn += 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 stallCo, Ci, kH, kW = kernel.shapebW_half = bW // 2 + 1bhw = bH * bW_halfW_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.shapeCo, _, kH, kW = kernel.shapeH_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 + 1bW_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 blockstotal_h = (n_h - 1) * L_h + bHtotal_w = (n_w - 1) * L_w + bWpad_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 unfoldx_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 linesreturn 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 = paramelse:- 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