Skip to content
KernelIndex
Search⌘K

submission 545541

krasnaya_66854 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_b200_cl_fast.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-545541?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.2ms
#16 of 28
2026-03-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4e354df3e89c01ae32ec44a2971cf820767951f72a7c04d27051d02d4eb098b4
license declaredunknown
license concludedunknown
authorskrasnaya_66854
imported2026-08-15

Kernel source

submission_b200_cl_fast.py255 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=False,
        benchmark=True,
        allow_tf32=False,
    ):
        output[...] = F.conv2d(input_tensor, kernel, stride=1, padding=0)
    return output


def conv2d_cudnn_exact(
    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=False,
        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_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=False,
        benchmark=True,
        allow_tf32=False,
    ):
        y = F.conv2d(x, w, stride=1, padding=0)
    output[...] = y.contiguous()
    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_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): ("fft_full", None),
    (1, 128, 256, 32): ("cudnn_cl", None),
}


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

⋯ 65 unchanged lines
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,
+ deterministic=False,
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:
+ prev = torch.are_deterministic_algorithms_enabled()
+ with torch.backends.cudnn.flags(
+ deterministic=True,
+ benchmark=False,
+ allow_tf32=False,
+ ):
+ torch.use_deterministic_algorithms(True)
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)
+ torch.use_deterministic_algorithms(prev)
return output
⋯ 2 unchanged lines
) -> 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,
+ deterministic=False,
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,
⋯ 117 unchanged lines
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":
⋯ 11 unchanged lines
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),
+ (2, 128, 256, 16): ("fft_full", None),
+ (1, 128, 256, 32): ("cudnn_cl", 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
+ return run_plan(data, PLAN)
scrolls · 151 diff lines total

Best evidence level for this revision: reported

JSON