Skip to content
KernelIndex
Search⌘K

submission 615496

KernelAgent · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

gpumode_submit_g5vokthf.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-615496?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
128.9ms
#18 of 35
2026-03-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:77e666fb97e7740cbc988a7eeb64826600a170786608a97421043e82367970c3
license declaredunknown
license concludedunknown
authorsKernelAgent
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

mmaacc = tl.dot(a, b, acc, allow_tf32=False)
tile-k = 16BLOCK_K = 16
tile-m = 64BLOCK_M = 64

Kernel source

gpumode_submit_g5vokthf.py135 lines
import torch
import triton
import triton.language as tl


@triton.jit
def _conv2d_implicit_gemm_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,
    total_out_pixels,
    reduction_size,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """
    Fused 2D convolution via implicit GEMM.
    M = out_h*out_w (output spatial), N = out_channels, K = in_channels*kernel_h*kernel_w
    A[M,K] = im2col(input), B[K,N] = weight reshaped, C[M,N] = output
    """
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    pid_b = tl.program_id(2)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    m_mask = offs_m < total_out_pixels
    n_mask = offs_n < out_channels

    oh = offs_m // out_w
    ow = offs_m % out_w

    acc = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)

    kh_kw = kernel_h * kernel_w
    base_in = input_ptr + pid_b * stride_in_b

    for k_start in range(0, reduction_size, BLOCK_K):
        offs_k = k_start + tl.arange(0, BLOCK_K)
        k_mask = offs_k < reduction_size

        ic = offs_k // kh_kw
        rem = offs_k % kh_kw
        kh = rem // kernel_w
        kw = rem % kernel_w

        # A tile [BLOCK_M, BLOCK_K]: implicit im2col
        ih = oh[:, None] + kh[None, :]
        iw = ow[:, None] + kw[None, :]
        in_ptrs = base_in + ic[None, :] * stride_in_c + ih * stride_in_h + iw * stride_in_w
        a_mask = m_mask[:, None] & k_mask[None, :]
        a = tl.load(in_ptrs, mask=a_mask, other=0.0)

        # B tile [BLOCK_K, BLOCK_N]: weight[oc, ic, kh, kw]
        w_ptrs = (weight_ptr
                  + offs_n[None, :] * stride_w_oc
                  + ic[:, None] * stride_w_ic
                  + kh[:, None] * stride_w_kh
                  + kw[:, None] * stride_w_kw)
        b_mask = k_mask[:, None] & n_mask[None, :]
        b = tl.load(w_ptrs, mask=b_mask, other=0.0)

        acc = tl.dot(a, b, acc, allow_tf32=False)

    # Store output
    out_ptrs = (output_ptr + pid_b * stride_out_b
                + offs_n[None, :] * stride_out_c
                + oh[:, None] * stride_out_h
                + ow[:, None] * stride_out_w)
    out_mask = m_mask[:, None] & n_mask[None, :]
    tl.store(out_ptrs, acc, mask=out_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

    if output_tensor is None:
        output_tensor = torch.empty(
            (batch, out_channels, out_h, out_w),
            device=input_tensor.device, dtype=input_tensor.dtype,
        )

    total_out_pixels = out_h * out_w
    reduction_size = in_channels * kernel_h * kernel_w

    BLOCK_M = 64
    BLOCK_N = min(64, triton.next_power_of_2(out_channels))
    BLOCK_K = min(32, triton.next_power_of_2(reduction_size))
    if BLOCK_K < 16:
        BLOCK_K = 16

    grid = (triton.cdiv(total_out_pixels, BLOCK_M),
            triton.cdiv(out_channels, BLOCK_N),
            batch)

    _conv2d_implicit_gemm_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),
        total_out_pixels,
        reduction_size,
        BLOCK_M=BLOCK_M,
        BLOCK_N=BLOCK_N,
        BLOCK_K=BLOCK_K,
    )
    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 · 135 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 597567.

⋯ 3 unchanged lines
@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,
+ def _conv2d_implicit_gemm_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,
+ total_out_pixels,
+ reduction_size,
+ BLOCK_M: tl.constexpr,
+ BLOCK_N: tl.constexpr,
+ BLOCK_K: 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.
+ Fused 2D convolution via implicit GEMM.
+ M = out_h*out_w (output spatial), N = out_channels, K = in_channels*kernel_h*kernel_w
+ A[M,K] = im2col(input), B[K,N] = weight reshaped, C[M,N] = output
"""
- 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
+ pid_m = tl.program_id(0)
+ pid_n = tl.program_id(1)
+ pid_b = tl.program_id(2)
- # Tile starts
- ow_start = pid_w * BLOCK_W
- oh_start = pid_h * BLOCK_H
+ offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
- # 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
+ m_mask = offs_m < total_out_pixels
+ n_mask = offs_n < out_channels
- # 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)
+ oh = offs_m // out_w
+ ow = offs_m % out_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, :]
+ acc = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)
- # Base offsets for this batch element
- n_off_inp = n_id * sN
- n_off_out = n_id * soN
+ kh_kw = kernel_h * kernel_w
+ base_in = input_ptr + pid_b * stride_in_b
- # Accumulator in FP32: [BLOCK_OC, BLOCK_H, BLOCK_W]
- acc = tl.zeros((BLOCK_OC, BLOCK_H, BLOCK_W), dtype=tl.float32)
+ for k_start in range(0, reduction_size, BLOCK_K):
+ offs_k = k_start + tl.arange(0, BLOCK_K)
+ k_mask = offs_k < reduction_size
- # 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
+ ic = offs_k // kh_kw
+ rem = offs_k % kh_kw
+ kh = rem // kernel_w
+ kw = rem % kernel_w
- # 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)
+ # A tile [BLOCK_M, BLOCK_K]: implicit im2col
+ ih = oh[:, None] + kh[None, :]
+ iw = ow[:, None] + kw[None, :]
+ in_ptrs = base_in + ic[None, :] * stride_in_c + ih * stride_in_h + iw * stride_in_w
+ a_mask = m_mask[:, None] & k_mask[None, :]
+ a = tl.load(in_ptrs, mask=a_mask, 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)
+ # B tile [BLOCK_K, BLOCK_N]: weight[oc, ic, kh, kw]
+ w_ptrs = (weight_ptr
+ + offs_n[None, :] * stride_w_oc
+ + ic[:, None] * stride_w_ic
+ + kh[:, None] * stride_w_kh
+ + kw[:, None] * stride_w_kw)
+ b_mask = k_mask[:, None] & n_mask[None, :]
+ b = tl.load(w_ptrs, mask=b_mask, other=0.0)
- # FMA: broadcast multiply-accumulate
- # acc[oc, oh, ow] += w_vec[oc] * x_patch[oh, ow]
- acc += w_vec[:, None, None] * x_patch[None, :, :]
+ acc = tl.dot(a, b, acc, allow_tf32=False)
- # 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)
+ # Store output
+ out_ptrs = (output_ptr + pid_b * stride_out_b
+ + offs_n[None, :] * stride_out_c
+ + oh[:, None] * stride_out_h
+ + ow[:, None] * stride_out_w)
+ out_mask = m_mask[:, None] & n_mask[None, :]
+ tl.store(out_ptrs, acc, mask=out_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]
+ 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
- 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"
+ output_tensor = torch.empty(
+ (batch, out_channels, out_h, out_w),
+ device=input_tensor.device, dtype=input_tensor.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()
+ total_out_pixels = out_h * out_w
+ reduction_size = in_channels * kernel_h * kernel_w
- # Choose block sizes (powers of two for good performance)
- BLOCK_OC = 16
- BLOCK_H = 8
- BLOCK_W = 64
+ BLOCK_M = 64
+ BLOCK_N = min(64, triton.next_power_of_2(out_channels))
+ BLOCK_K = min(32, triton.next_power_of_2(reduction_size))
+ if BLOCK_K < 16:
+ BLOCK_K = 16
- # 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
- )
+ grid = (triton.cdiv(total_out_pixels, BLOCK_M),
+ triton.cdiv(out_channels, BLOCK_N),
+ batch)
- # 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,
+ _conv2d_implicit_gemm_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),
+ total_out_pixels,
+ reduction_size,
+ BLOCK_M=BLOCK_M,
+ BLOCK_N=BLOCK_N,
+ BLOCK_K=BLOCK_K,
)
-
return output_tensor
import inspect
scrolls · 258 diff lines total

Best evidence level for this revision: reported

JSON