Skip to content
KernelIndex
Search⌘K

submission 597567

KernelAgent · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

gpumode_submit_oj47ymbh.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-597567?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
2D convolutionsuite of 5 cases
NVIDIA H100
131.5ms
#20 of 35
2026-03-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5022d75f21a4dcf4e538d5ce96cc4bb366dbe1248ef18d30c5f483a74c3d8d06
license declaredunknown
license concludedunknown
authorsKernelAgent
imported2026-08-15

Kernel source

gpumode_submit_oj47ymbh.py179 lines
import torch
import triton
import triton.language as tl


@triton.jit
def _conv2d_nchw_kernel(
    inp_ptr,            # *T[N, CI, H, W]
    w_ptr,              # *T[CO, CI, KH, KW]
    out_ptr,            # *T[N, CO, H_out, W_out]
    N, CI, CO, H, W, KH, KW,
    H_out, W_out,
    sN, sC, sH, sW,     # input strides
    sw_oc, sw_ic, sw_kh, sw_kw,  # weight strides
    soN, soC, soH, soW,          # output strides
    BLOCK_OC: tl.constexpr,
    BLOCK_H: tl.constexpr,
    BLOCK_W: tl.constexpr,
):
    """
    Compute 2D convolution (NCHW, stride=1, no padding) in Triton.

    Each program computes a tile of the output tensor with shape:
      [BLOCK_OC, BLOCK_H, BLOCK_W] across (oc, oh, ow).
    Reductions over input channels and kernel taps are done in FP32 for stability.
    """
    pid_w = tl.program_id(axis=0)  # tile along output width
    pid_h = tl.program_id(axis=1)  # tile along output height
    pid_noc = tl.program_id(axis=2)  # fused batch and oc-block id

    # Tile starts
    ow_start = pid_w * BLOCK_W
    oh_start = pid_h * BLOCK_H

    # Derive n and oc-block ids from fused pid
    num_oc_blocks = tl.cdiv(CO, BLOCK_OC)
    n_id = pid_noc // num_oc_blocks
    oc_block_id = pid_noc % num_oc_blocks

    # Offsets for this tile
    offs_oc = oc_block_id * BLOCK_OC + tl.arange(0, BLOCK_OC)
    offs_oh = oh_start + tl.arange(0, BLOCK_H)
    offs_ow = ow_start + tl.arange(0, BLOCK_W)

    # Masks for boundaries
    oc_mask = offs_oc < CO
    oh_mask = offs_oh < H_out
    ow_mask = offs_ow < W_out
    spatial_mask = oh_mask[:, None] & ow_mask[None, :]

    # Base offsets for this batch element
    n_off_inp = n_id * sN
    n_off_out = n_id * soN

    # Accumulator in FP32: [BLOCK_OC, BLOCK_H, BLOCK_W]
    acc = tl.zeros((BLOCK_OC, BLOCK_H, BLOCK_W), dtype=tl.float32)

    # Iterate over input channels and kernel taps
    for ic in range(0, CI):
        ic_off_inp = ic * sC
        ic_off_w = ic * sw_ic
        # For each kernel row/col
        for ky in range(0, KH):
            ky_off_inp = (offs_oh + ky) * sH
            ky_off_w = ky * sw_kh
            for kx in range(0, KW):
                kx_off_inp = (offs_ow + kx) * sW
                kx_off_w = kx * sw_kw

                # Load input patch: [BLOCK_H, BLOCK_W]
                inp_ptrs = inp_ptr + n_off_inp + ic_off_inp + ky_off_inp[:, None] + kx_off_inp[None, :]
                x_patch = tl.load(inp_ptrs, mask=spatial_mask, other=0.0).to(tl.float32)

                # Load weights for [BLOCK_OC] at (oc, ic, ky, kx)
                w_ptrs = w_ptr + offs_oc * sw_oc + ic_off_w + ky_off_w + kx_off_w
                w_vec = tl.load(w_ptrs, mask=oc_mask, other=0.0).to(tl.float32)

                # FMA: broadcast multiply-accumulate
                # acc[oc, oh, ow] += w_vec[oc] * x_patch[oh, ow]
                acc += w_vec[:, None, None] * x_patch[None, :, :]

    # Store results to output (cast to output dtype)
    out_vals = acc.to(out_ptr.dtype.element_ty)
    out_ptrs = (
        out_ptr
        + n_off_out
        + offs_oc[:, None, None] * soC
        + offs_oh[None, :, None] * soH
        + offs_ow[None, None, :] * soW
    )
    store_mask = oc_mask[:, None, None] & spatial_mask[None, :, :]
    tl.store(out_ptrs, out_vals, mask=store_mask)


def kernel_function(input_tensor, kernel, output_tensor=None):
    """
    Triton 2D convolution (NCHW) wrapper.
    Implements stride=1, no padding: out[n, oc, oh, ow] = sum_{ic, ky, kx} x[n, ic, oh+ky, ow+kx] * w[oc, ic, ky, kx]

    Fusion discussion:
    - This problem requires only a plain conv2d. There are no additional stages
      (bias, activation, pooling) in the pipeline to fuse. Therefore, we run a
      single Triton kernel that performs the full convolution and reduction in
      one pass, with FP32 accumulation and direct store to the output dtype.

    Args:
      input_tensor:  torch.Tensor [N, CI, H, W], dtype: float32 or bfloat16, on CUDA
      kernel:        torch.Tensor [CO, CI, KH, KW], same dtype/device as input_tensor
      output_tensor: optional preallocated tensor [N, CO, H-KH+1, W-KW+1]; if None, it's allocated

    Returns:
      torch.Tensor with shape [N, CO, H_out, W_out], same dtype/device as inputs.
    """
    # Basic validation and setup (no math here; all compute happens in the Triton kernel)
    assert isinstance(input_tensor, torch.Tensor) and isinstance(kernel, torch.Tensor), "Inputs must be tensors"
    assert input_tensor.device.type == "cuda" and kernel.device.type == "cuda", "Tensors must be on CUDA device"
    assert input_tensor.dtype == kernel.dtype, "Input and kernel dtypes must match"
    assert input_tensor.ndim == 4 and kernel.ndim == 4, "Expected NCHW input and OIHW kernel"

    N, CI, H, W = input_tensor.shape
    CO, CI_k, KH, KW = kernel.shape
    assert CI_k == CI, "kernel in_channels must match input channels"
    # No padding, stride=1 -> output sizes
    H_out = H - KH + 1
    W_out = W - KW + 1
    assert H_out > 0 and W_out > 0, "Kernel size must be <= input spatial size (no padding)."

    # Allocate output if not provided
    if output_tensor is None:
        output_tensor = torch.empty((N, CO, H_out, W_out), device=input_tensor.device, dtype=input_tensor.dtype)
    else:
        assert output_tensor.shape == (N, CO, H_out, W_out), "Output tensor has incorrect shape"
        assert output_tensor.device == input_tensor.device and output_tensor.dtype == input_tensor.dtype, \
            "Output tensor must match input device/dtype"

    # Strides (in elements)
    sN, sC, sH, sW_ = input_tensor.stride()
    sw_oc, sw_ic, sw_kh, sw_kw = kernel.stride()
    soN, soC, soH, soW = output_tensor.stride()

    # Choose block sizes (powers of two for good performance)
    BLOCK_OC = 16
    BLOCK_H = 8
    BLOCK_W = 64

    # Grid dimensions
    grid = (
        triton.cdiv(W_out, BLOCK_W),               # tiles along width
        triton.cdiv(H_out, BLOCK_H),               # tiles along height
        N * triton.cdiv(CO, BLOCK_OC),             # fused batch and oc blocks
    )

    # Launch kernel
    _conv2d_nchw_kernel[grid](
        input_tensor, kernel, output_tensor,
        N, CI, CO, H, W, KH, KW,
        H_out, W_out,
        sN, sC, sH, sW_,
        sw_oc, sw_ic, sw_kh, sw_kw,
        soN, soC, soH, soW,
        BLOCK_OC=BLOCK_OC,
        BLOCK_H=BLOCK_H,
        BLOCK_W=BLOCK_W,
    )

    return output_tensor

import inspect
def custom_kernel(input):
    sig = inspect.signature(kernel_function)
    num_params = len(sig.parameters)
    if len(input) == num_params:
        return kernel_function(*input)
    return kernel_function(input)

import os
if os.environ.get("CUBLAS_WORKSPACE_CONFIG", "") not in (":4096:8", ":16:8"):
    os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
scrolls · 179 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 597516.

⋯ 3 unchanged lines
@triton.jit
- def _conv2d_kernel(
- input_ptr, weight_ptr, output_ptr,
- in_channels, out_channels,
- in_h, in_w, out_h, out_w,
- kernel_h, kernel_w,
- stride_in_b, stride_in_c, stride_in_h, stride_in_w,
- stride_w_oc, stride_w_ic, stride_w_kh, stride_w_kw,
- stride_out_b, stride_out_c, stride_out_h, stride_out_w,
- K, # in_channels * kernel_h * kernel_w
- M, # out_h * out_w
- BLOCK_M: tl.constexpr,
- BLOCK_N: tl.constexpr,
- BLOCK_K: tl.constexpr,
+ def _conv2d_nchw_kernel(
+ inp_ptr, # *T[N, CI, H, W]
+ w_ptr, # *T[CO, CI, KH, KW]
+ out_ptr, # *T[N, CO, H_out, W_out]
+ N, CI, CO, H, W, KH, KW,
+ H_out, W_out,
+ sN, sC, sH, sW, # input strides
+ sw_oc, sw_ic, sw_kh, sw_kw, # weight strides
+ soN, soC, soH, soW, # output strides
+ BLOCK_OC: tl.constexpr,
+ BLOCK_H: tl.constexpr,
+ BLOCK_W: tl.constexpr,
):
- pid_m = tl.program_id(0)
- pid_n = tl.program_id(1)
- b = tl.program_id(2)
+ """
+ Compute 2D convolution (NCHW, stride=1, no padding) in Triton.
- offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
- offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
+ Each program computes a tile of the output tensor with shape:
+ [BLOCK_OC, BLOCK_H, BLOCK_W] across (oc, oh, ow).
+ Reductions over input channels and kernel taps are done in FP32 for stability.
+ """
+ pid_w = tl.program_id(axis=0) # tile along output width
+ pid_h = tl.program_id(axis=1) # tile along output height
+ pid_noc = tl.program_id(axis=2) # fused batch and oc-block id
- mask_m = offs_m < M
- mask_n = offs_n < out_channels
+ # Tile starts
+ ow_start = pid_w * BLOCK_W
+ oh_start = pid_h * BLOCK_H
- oh = offs_m // out_w
- ow = offs_m % out_w
+ # Derive n and oc-block ids from fused pid
+ num_oc_blocks = tl.cdiv(CO, BLOCK_OC)
+ n_id = pid_noc // num_oc_blocks
+ oc_block_id = pid_noc % num_oc_blocks
- acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ # Offsets for this tile
+ offs_oc = oc_block_id * BLOCK_OC + tl.arange(0, BLOCK_OC)
+ offs_oh = oh_start + tl.arange(0, BLOCK_H)
+ offs_ow = ow_start + tl.arange(0, BLOCK_W)
- kh_kw = kernel_h * kernel_w
+ # Masks for boundaries
+ oc_mask = offs_oc < CO
+ oh_mask = offs_oh < H_out
+ ow_mask = offs_ow < W_out
+ spatial_mask = oh_mask[:, None] & ow_mask[None, :]
- for k_start in range(0, K, BLOCK_K):
- offs_k = k_start + tl.arange(0, BLOCK_K)
- mask_k = offs_k < K
+ # Base offsets for this batch element
+ n_off_inp = n_id * sN
+ n_off_out = n_id * soN
- ic = offs_k // kh_kw
- rem = offs_k % kh_kw
- kh = rem // kernel_w
- kw = rem % kernel_w
+ # Accumulator in FP32: [BLOCK_OC, BLOCK_H, BLOCK_W]
+ acc = tl.zeros((BLOCK_OC, BLOCK_H, BLOCK_W), dtype=tl.float32)
- ih = oh[:, None] + kh[None, :]
- iw = ow[:, None] + kw[None, :]
+ # Iterate over input channels and kernel taps
+ for ic in range(0, CI):
+ ic_off_inp = ic * sC
+ ic_off_w = ic * sw_ic
+ # For each kernel row/col
+ for ky in range(0, KH):
+ ky_off_inp = (offs_oh + ky) * sH
+ ky_off_w = ky * sw_kh
+ for kx in range(0, KW):
+ kx_off_inp = (offs_ow + kx) * sW
+ kx_off_w = kx * sw_kw
- in_offset = (b * stride_in_b +
- ic[None, :] * stride_in_c +
- ih * stride_in_h +
- iw * stride_in_w)
- a_tile = tl.load(input_ptr + in_offset,
- mask=mask_m[:, None] & mask_k[None, :], other=0.0)
+ # Load input patch: [BLOCK_H, BLOCK_W]
+ inp_ptrs = inp_ptr + n_off_inp + ic_off_inp + ky_off_inp[:, None] + kx_off_inp[None, :]
+ x_patch = tl.load(inp_ptrs, mask=spatial_mask, other=0.0).to(tl.float32)
- w_offset = (offs_n[:, None] * stride_w_oc +
- ic[None, :] * stride_w_ic +
- kh[None, :] * stride_w_kh +
- kw[None, :] * stride_w_kw)
- w_tile = tl.load(weight_ptr + w_offset,
- mask=mask_n[:, None] & mask_k[None, :], other=0.0)
+ # Load weights for [BLOCK_OC] at (oc, ic, ky, kx)
+ w_ptrs = w_ptr + offs_oc * sw_oc + ic_off_w + ky_off_w + kx_off_w
+ w_vec = tl.load(w_ptrs, mask=oc_mask, other=0.0).to(tl.float32)
- acc += tl.dot(a_tile, tl.trans(w_tile), allow_tf32=False)
+ # FMA: broadcast multiply-accumulate
+ # acc[oc, oh, ow] += w_vec[oc] * x_patch[oh, ow]
+ acc += w_vec[:, None, None] * x_patch[None, :, :]
- out_offset = (b * stride_out_b +
- offs_n[None, :] * stride_out_c +
- oh[:, None] * stride_out_h +
- ow[:, None] * stride_out_w)
- out_mask = mask_m[:, None] & mask_n[None, :]
- tl.store(output_ptr + out_offset, acc, mask=out_mask)
+ # Store results to output (cast to output dtype)
+ out_vals = acc.to(out_ptr.dtype.element_ty)
+ out_ptrs = (
+ out_ptr
+ + n_off_out
+ + offs_oc[:, None, None] * soC
+ + offs_oh[None, :, None] * soH
+ + offs_ow[None, None, :] * soW
+ )
+ store_mask = oc_mask[:, None, None] & spatial_mask[None, :, :]
+ tl.store(out_ptrs, out_vals, mask=store_mask)
- def kernel_function(input_tensor, kernel_weights, output_tensor=None):
- batch, in_channels, in_h, in_w = input_tensor.shape
- out_channels, _, kernel_h, kernel_w = kernel_weights.shape
- out_h = in_h - kernel_h + 1
- out_w = in_w - kernel_w + 1
+ def kernel_function(input_tensor, kernel, output_tensor=None):
+ """
+ Triton 2D convolution (NCHW) wrapper.
+ Implements stride=1, no padding: out[n, oc, oh, ow] = sum_{ic, ky, kx} x[n, ic, oh+ky, ow+kx] * w[oc, ic, ky, kx]
+ Fusion discussion:
+ - This problem requires only a plain conv2d. There are no additional stages
+ (bias, activation, pooling) in the pipeline to fuse. Therefore, we run a
+ single Triton kernel that performs the full convolution and reduction in
+ one pass, with FP32 accumulation and direct store to the output dtype.
+
+ Args:
+ input_tensor: torch.Tensor [N, CI, H, W], dtype: float32 or bfloat16, on CUDA
+ kernel: torch.Tensor [CO, CI, KH, KW], same dtype/device as input_tensor
+ output_tensor: optional preallocated tensor [N, CO, H-KH+1, W-KW+1]; if None, it's allocated
+
+ Returns:
+ torch.Tensor with shape [N, CO, H_out, W_out], same dtype/device as inputs.
+ """
+ # Basic validation and setup (no math here; all compute happens in the Triton kernel)
+ assert isinstance(input_tensor, torch.Tensor) and isinstance(kernel, torch.Tensor), "Inputs must be tensors"
+ assert input_tensor.device.type == "cuda" and kernel.device.type == "cuda", "Tensors must be on CUDA device"
+ assert input_tensor.dtype == kernel.dtype, "Input and kernel dtypes must match"
+ assert input_tensor.ndim == 4 and kernel.ndim == 4, "Expected NCHW input and OIHW kernel"
+
+ N, CI, H, W = input_tensor.shape
+ CO, CI_k, KH, KW = kernel.shape
+ assert CI_k == CI, "kernel in_channels must match input channels"
+ # No padding, stride=1 -> output sizes
+ H_out = H - KH + 1
+ W_out = W - KW + 1
+ assert H_out > 0 and W_out > 0, "Kernel size must be <= input spatial size (no padding)."
+
+ # Allocate output if not provided
if output_tensor is None:
- output_tensor = torch.empty(
- (batch, out_channels, out_h, out_w),
- device=input_tensor.device, dtype=input_tensor.dtype)
+ output_tensor = torch.empty((N, CO, H_out, W_out), device=input_tensor.device, dtype=input_tensor.dtype)
+ else:
+ assert output_tensor.shape == (N, CO, H_out, W_out), "Output tensor has incorrect shape"
+ assert output_tensor.device == input_tensor.device and output_tensor.dtype == input_tensor.dtype, \
+ "Output tensor must match input device/dtype"
- M = out_h * out_w
- K = in_channels * kernel_h * kernel_w
+ # Strides (in elements)
+ sN, sC, sH, sW_ = input_tensor.stride()
+ sw_oc, sw_ic, sw_kh, sw_kw = kernel.stride()
+ soN, soC, soH, soW = output_tensor.stride()
- BLOCK_M = 64
- BLOCK_N = 64
- BLOCK_K = 32
+ # Choose block sizes (powers of two for good performance)
+ BLOCK_OC = 16
+ BLOCK_H = 8
+ BLOCK_W = 64
- grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(out_channels, BLOCK_N), batch)
+ # Grid dimensions
+ grid = (
+ triton.cdiv(W_out, BLOCK_W), # tiles along width
+ triton.cdiv(H_out, BLOCK_H), # tiles along height
+ N * triton.cdiv(CO, BLOCK_OC), # fused batch and oc blocks
+ )
- _conv2d_kernel[grid](
- input_tensor, kernel_weights, output_tensor,
- in_channels, out_channels,
- in_h, in_w, out_h, out_w,
- kernel_h, kernel_w,
- input_tensor.stride(0), input_tensor.stride(1),
- input_tensor.stride(2), input_tensor.stride(3),
- kernel_weights.stride(0), kernel_weights.stride(1),
- kernel_weights.stride(2), kernel_weights.stride(3),
- output_tensor.stride(0), output_tensor.stride(1),
- output_tensor.stride(2), output_tensor.stride(3),
- K, M,
- BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
+ # Launch kernel
+ _conv2d_nchw_kernel[grid](
+ input_tensor, kernel, output_tensor,
+ N, CI, CO, H, W, KH, KW,
+ H_out, W_out,
+ sN, sC, sH, sW_,
+ sw_oc, sw_ic, sw_kh, sw_kw,
+ soN, soC, soH, soW,
+ BLOCK_OC=BLOCK_OC,
+ BLOCK_H=BLOCK_H,
+ BLOCK_W=BLOCK_W,
)
+
return output_tensor
import inspect
scrolls · 245 diff lines total

Best evidence level for this revision: reported

JSON