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
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.
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 providedif 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_tensorimport inspect
scrolls · 258 diff lines total
Best evidence level for this revision: reported
JSON