submission 545518
krasnaya_66854 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 330 lines, June 9 Researcher Reciprocity License v1.0.
submission_h100_graph.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-545518?include=source"interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
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:a93e065e22243f1fc3fee108cd51f14323fa2e900ad263b17fba85c4eec71e1e
license declaredunknown
license concludedunknown
authorskrasnaya_66854
imported2026-08-15
Kernel source
submission_h100_graph.py330 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 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_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 == "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): ("cudnn_cl", None),
(1, 128, 256, 32): ("cudnn_exact", None),
}
_GRAPH_WARM_KEYS: set[tuple[int, int, int, int, int, int, int]] = set()
_GRAPH_CACHE: dict[tuple[int, int, int, int, int, int, int], torch.cuda.CUDAGraph] = {}
_GRAPH_DISABLED: set[tuple[int, int, int, int]] = set()
def custom_kernel(data: input_t) -> output_t:
input_tensor, kernel, output = data
shape = shape_key(input_tensor, kernel)
if shape != (1, 128, 256, 32) or shape in _GRAPH_DISABLED:
return run_plan(data, PLAN)
key = (
shape[0],
shape[1],
shape[2],
shape[3],
input_tensor.data_ptr(),
kernel.data_ptr(),
output.data_ptr(),
)
graph = _GRAPH_CACHE.get(key)
if graph is not None:
graph.replay()
return output
if key not in _GRAPH_WARM_KEYS:
_GRAPH_WARM_KEYS.add(key)
return run_plan(data, PLAN)
graph = torch.cuda.CUDAGraph()
try:
with torch.cuda.graph(graph):
run_plan(data, PLAN)
except Exception:
_GRAPH_DISABLED.add(shape)
return run_plan(data, PLAN)
_GRAPH_CACHE[key] = graph
return output
scrolls · 330 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 543980.
⋯ 116 unchanged linesreturn 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,⋯ 20 unchanged linesreturn 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- 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:- 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,- )- 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 fft_conv2d_full(input_tensor: torch.Tensor, kernel: torch.Tensor, output: torch.Tensor) -> torch.Tensor:⋯ 91 unchanged linesreturn 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 16))if method == "fft_full":return fft_conv2d_full(input_tensor, kernel, output)⋯ 7 unchanged linesPLAN = {(4, 64, 128, 8): ("cudnn_cl", None),(4, 64, 128, 16): ("cudnn_cl", None),- (2, 128, 256, 16): ("cudnn", None),- (1, 128, 256, 32): ("cudnn_tiled_exact", 48),+ (2, 128, 256, 16): ("cudnn_cl", None),+ (1, 128, 256, 32): ("cudnn_exact", None),}+ _GRAPH_WARM_KEYS: set[tuple[int, int, int, int, int, int, int]] = set()+ _GRAPH_CACHE: dict[tuple[int, int, int, int, int, int, int], torch.cuda.CUDAGraph] = {}+ _GRAPH_DISABLED: set[tuple[int, int, int, int]] = set()++def custom_kernel(data: input_t) -> output_t:- return run_plan(data, PLAN)+ input_tensor, kernel, output = data+ shape = shape_key(input_tensor, kernel)+ if shape != (1, 128, 256, 32) or shape in _GRAPH_DISABLED:+ return run_plan(data, PLAN)++ key = (+ shape[0],+ shape[1],+ shape[2],+ shape[3],+ input_tensor.data_ptr(),+ kernel.data_ptr(),+ output.data_ptr(),+ )++ graph = _GRAPH_CACHE.get(key)+ if graph is not None:+ graph.replay()+ return output++ if key not in _GRAPH_WARM_KEYS:+ _GRAPH_WARM_KEYS.add(key)+ return run_plan(data, PLAN)++ graph = torch.cuda.CUDAGraph()+ try:+ with torch.cuda.graph(graph):+ run_plan(data, PLAN)+ except Exception:+ _GRAPH_DISABLED.add(shape)+ return run_plan(data, PLAN)++ _GRAPH_CACHE[key] = graph+ return output
scrolls · 138 diff lines total
Best evidence level for this revision: reported
JSON