submission 547141
krasnaya_66854 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 390 lines, June 9 Researcher Reciprocity License v1.0.
submission_b200_t152_fast.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-547141?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:298a57a6be9ccd6e302b1de0a79475d1774b52f05e311e1cd8d44e717c634f6a
license declaredunknown
license concludedunknown
authorskrasnaya_66854
imported2026-08-15
Kernel source
submission_b200_t152_fast.py390 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:
prev = torch.are_deterministic_algorithms_enabled()
with torch.backends.cudnn.flags(
deterministic=True,
benchmark=True,
allow_tf32=False,
):
torch.use_deterministic_algorithms(True)
output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
torch.use_deterministic_algorithms(prev)
return output
def conv2d_cudnn_exact(
input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor
) -> torch.Tensor:
prev_allow_tf32 = torch.backends.cudnn.allow_tf32
prev_deterministic = torch.backends.cudnn.deterministic
prev_benchmark = torch.backends.cudnn.benchmark
prev_algorithms = torch.are_deterministic_algorithms_enabled()
torch.backends.cudnn.allow_tf32 = False
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
torch.use_deterministic_algorithms(True)
try:
output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
finally:
torch.backends.cudnn.allow_tf32 = prev_allow_tf32
torch.backends.cudnn.deterministic = prev_deterministic
torch.backends.cudnn.benchmark = prev_benchmark
torch.use_deterministic_algorithms(prev_algorithms)
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)
prev = torch.are_deterministic_algorithms_enabled()
with torch.backends.cudnn.flags(
deterministic=True,
benchmark=True,
allow_tf32=False,
):
torch.use_deterministic_algorithms(True)
y = F.conv2d(x, w, stride=1, padding=0)
torch.use_deterministic_algorithms(prev)
output[...] = y.contiguous()
return output
def conv2d_cudnn_channels_last_exact(
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)
prev_allow_tf32 = torch.backends.cudnn.allow_tf32
prev_deterministic = torch.backends.cudnn.deterministic
prev_benchmark = torch.backends.cudnn.benchmark
prev_algorithms = torch.are_deterministic_algorithms_enabled()
torch.backends.cudnn.allow_tf32 = False
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
torch.use_deterministic_algorithms(True)
try:
output[...] = F.conv2d(x, w, stride=1, padding=0).contiguous()
finally:
torch.backends.cudnn.allow_tf32 = prev_allow_tf32
torch.backends.cudnn.deterministic = prev_deterministic
torch.backends.cudnn.benchmark = prev_benchmark
torch.use_deterministic_algorithms(prev_algorithms)
return output
def conv2d_cudnn_chunked(
input_tensor: torch.Tensor,
kernel: torch.Tensor,
output: torch.Tensor,
chunk_size: int,
exact: bool = False,
) -> torch.Tensor:
prev = torch.are_deterministic_algorithms_enabled()
with torch.backends.cudnn.flags(
deterministic=True,
benchmark=not exact,
allow_tf32=False,
):
torch.use_deterministic_algorithms(True)
for start in range(0, kernel.shape[0], chunk_size):
end = min(start + chunk_size, kernel.shape[0])
output[:, start:end] = F.conv2d(
input_tensor,
kernel[start:end],
stride=1,
padding=0,
)
torch.use_deterministic_algorithms(prev)
return output
def conv2d_cudnn_tiled_exact(
input_tensor: torch.Tensor,
kernel: torch.Tensor,
output: torch.Tensor,
tile_h: int,
) -> torch.Tensor:
_, _, kH, _ = kernel.shape
H_out = input_tensor.shape[2] - kH + 1
with torch.backends.cudnn.flags(
deterministic=False,
benchmark=True,
allow_tf32=False,
):
for out_start in range(0, H_out, tile_h):
out_end = min(out_start + tile_h, H_out)
in_end = out_end + kH - 1
output[:, :, out_start:out_end, :] = F.conv2d(
input_tensor[:, :, out_start:in_end, :],
kernel,
stride=1,
padding=0,
)
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 fft_conv2d_block_fp64(
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
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)
.double()
)
w_blocks = kernel.double().reshape(Co * Ci, kH, kW)
W_fft_raw = torch.fft.rfft2(w_blocks, s=(bH, bW))
W_hw = (
W_fft_raw.view(Co, Ci, bH, bW_half)
.conj()
.permute(2, 3, 0, 1)
.reshape(bhw, Co, Ci)
.contiguous()
)
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].to(output.dtype)
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_exact":
return conv2d_cudnn_exact(input_tensor, kernel, output)
if method == "cudnn_cl":
return conv2d_cudnn_channels_last(input_tensor, kernel, output)
if method == "cudnn_cl_exact":
return conv2d_cudnn_channels_last_exact(input_tensor, kernel, output)
if method == "cudnn_chunked":
return conv2d_cudnn_chunked(input_tensor, kernel, output, int(param or 16))
if method == "cudnn_chunked_exact":
return conv2d_cudnn_chunked(input_tensor, kernel, output, int(param or 16), exact=True)
if method == "cudnn_tiled_exact":
return conv2d_cudnn_tiled_exact(input_tensor, kernel, output, int(param or 32))
if method == "fft_full":
return fft_conv2d_full(input_tensor, kernel, output)
if method == "fft_block_fp64":
if isinstance(param, tuple):
bH, bW = param
else:
bH = bW = param or 64
return fft_conv2d_block_fp64(input_tensor, kernel, output, bH, bW)
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_tiled_exact", 152),
}
def custom_kernel(data: input_t) -> output_t:
return run_plan(data, PLAN)
scrolls · 390 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 546126.
⋯ 380 unchanged lines(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_tiled_exact", 144),+ (1, 128, 256, 32): ("cudnn_tiled_exact", 152),}
Best evidence level for this revision: reported
JSON