Skip to content
KernelIndex
Search⌘K

submission 513456

KernelAgent · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

conv2d_v2_H100_gpt-5-2_ka_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-513456?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
180.9ms
#27 of 35
2026-03-06

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

num-warps = 4num_warps=4,

Kernel source

conv2d_v2_H100_gpt-5-2_ka_submission.py197 lines
import os
import sys
from typing import Tuple, TypeVar, TypedDict

import torch
import torch.nn.functional as F
import triton
import triton.language as tl

# ----------------------------
# Original reference utilities
# ----------------------------

def make_match_reference(reference: callable, **kwargs):
    def wrapped(data, output):
        return match_reference(data, output, reference=reference, **kwargs)
    return wrapped

def match_reference(data, output, reference, rtol=1e-3, atol=1e-3):
    ref = reference(data)
    ok = torch.allclose(output, ref, rtol=rtol, atol=atol)
    if not ok:
        max_abs = (output - ref).abs().max().item()
        max_rel = ((output - ref).abs() / (ref.abs() + 1e-12)).max().item()
        raise AssertionError(f"Mismatch: max_abs={max_abs} max_rel={max_rel}")
    return True

class DeterministicContext:
    def __init__(self):
        self.allow_tf32 = None
        self.deterministic = None
        self.cublas = None

    def __enter__(self):
        self.cublas = os.environ.get("CUBLAS_WORKSPACE_CONFIG", "")
        self.allow_tf32 = torch.backends.cudnn.allow_tf32
        self.deterministic = torch.backends.cudnn.deterministic
        torch.backends.cudnn.allow_tf32 = False
        torch.backends.cudnn.deterministic = True
        torch.use_deterministic_algorithms(True)
        return self

    def __exit__(self, exc_type, exc_value, traceback):
        torch.backends.cudnn.allow_tf32 = self.allow_tf32
        torch.backends.cudnn.deterministic = self.deterministic
        torch.use_deterministic_algorithms(False)
        os.environ["CUBLAS_WORKSPACE_CONFIG"] = self.cublas

input_t = TypeVar("input_t", bound=Tuple[torch.Tensor, torch.Tensor, torch.Tensor])
output_t = TypeVar("output_t", bound=torch.Tensor)

class TestSpec(TypedDict):
    size: int
    kernelsize: int
    channels: int
    batch: int
    seed: int

def ref_kernel(data: input_t) -> output_t:
    with DeterministicContext():
        input_tensor, kernel, output = data
        return F.conv2d(input_tensor, kernel, stride=1, padding=0)

def generate_input(size: int, kernelsize: int, channels: int, batch: int, seed: int) -> input_t:
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)
    x = torch.randn(batch, channels, size, size, device="cuda", dtype=torch.float32, generator=gen).contiguous()
    w = torch.randn(channels, channels, kernelsize, kernelsize, device="cuda", dtype=torch.float32, generator=gen).contiguous()
    y = torch.empty(batch, channels, size - kernelsize + 1, size - kernelsize + 1, device="cuda", dtype=torch.float32)
    return x, w, y

check_implementation = make_match_reference(ref_kernel, rtol=1e-3, atol=1e-3)

# ----------------------------
# Triton implementation
# ----------------------------

@triton.jit
def _conv2d_nchw_fwd_kernel(
    x_ptr, w_ptr, y_ptr,
    N: tl.constexpr, C: tl.constexpr, H: tl.constexpr, W: tl.constexpr,
    KH: tl.constexpr, KW: tl.constexpr,
    OH: tl.constexpr, OW: tl.constexpr,
    stride_xn: tl.constexpr, stride_xc: tl.constexpr, stride_xh: tl.constexpr, stride_xw: tl.constexpr,
    stride_wn: tl.constexpr, stride_wc: tl.constexpr, stride_wh: tl.constexpr, stride_ww: tl.constexpr,
    stride_yn: tl.constexpr, stride_yc: tl.constexpr, stride_yh: tl.constexpr, stride_yw: tl.constexpr,
    BLOCK_P: tl.constexpr,   # pixels (OH*OW) per program
):
    pid_n = tl.program_id(0)
    pid_co = tl.program_id(1)
    pid_p = tl.program_id(2)

    p0 = pid_p * BLOCK_P + tl.arange(0, BLOCK_P)
    mask_p = p0 < (OH * OW)

    oh = p0 // OW
    ow = p0 - oh * OW

    acc = tl.zeros((BLOCK_P,), dtype=tl.float32)

    # Loop over C and KH*KW
    for ci in range(0, C):
        # base pointer for x at this (n, ci)
        x_base_nc = x_ptr + pid_n * stride_xn + ci * stride_xc
        # base pointer for w at this (co, ci)
        w_base_coci = w_ptr + pid_co * stride_wn + ci * stride_wc

        for kh in range(0, KH):
            for kw in range(0, KW):
                x_ptrs = x_base_nc + (oh + kh) * stride_xh + (ow + kw) * stride_xw
                w_val = tl.load(w_base_coci + kh * stride_wh + kw * stride_ww)
                x_val = tl.load(x_ptrs, mask=mask_p, other=0.0)
                acc += x_val * w_val

    y_ptrs = y_ptr + pid_n * stride_yn + pid_co * stride_yc + oh * stride_yh + ow * stride_yw
    tl.store(y_ptrs, acc, mask=mask_p)

def kernel_function(input_tensor: torch.Tensor, kernel: torch.Tensor, output_tensor: torch.Tensor) -> torch.Tensor:
    # Validate (no torch math ops)
    if input_tensor.device.type != "cuda" or kernel.device.type != "cuda" or output_tensor.device.type != "cuda":
        raise ValueError("All tensors must be on CUDA")
    if input_tensor.dtype != torch.float32 or kernel.dtype != torch.float32 or output_tensor.dtype != torch.float32:
        raise ValueError("Only float32 supported")
    if not input_tensor.is_contiguous() or not kernel.is_contiguous() or not output_tensor.is_contiguous():
        raise ValueError("All tensors must be contiguous")
    if input_tensor.ndim != 4 or kernel.ndim != 4 or output_tensor.ndim != 4:
        raise ValueError("Expected NCHW tensors")
    N, C, H, W = input_tensor.shape
    Cout, Cin, KH, KW = kernel.shape
    if Cout != C or Cin != C:
        raise ValueError("This problem expects kernel shape [C, C, KH, KW]")
    OH = H - KH + 1
    OW = W - KW + 1
    if output_tensor.shape != (N, Cout, OH, OW):
        raise ValueError("output_tensor has incorrect shape")

    # Launch: 3D grid (N, Cout, blocks over OH*OW)
    BLOCK_P = 256

    def grid(meta):
        return (N, Cout, triton.cdiv(OH * OW, meta["BLOCK_P"]))

    _conv2d_nchw_fwd_kernel[grid](
        input_tensor, kernel, output_tensor,
        N=N, C=C, H=H, W=W, KH=KH, KW=KW, OH=OH, OW=OW,
        stride_xn=input_tensor.stride(0), stride_xc=input_tensor.stride(1),
        stride_xh=input_tensor.stride(2), stride_xw=input_tensor.stride(3),
        stride_wn=kernel.stride(0), stride_wc=kernel.stride(1),
        stride_wh=kernel.stride(2), stride_ww=kernel.stride(3),
        stride_yn=output_tensor.stride(0), stride_yc=output_tensor.stride(1),
        stride_yh=output_tensor.stride(2), stride_yw=output_tensor.stride(3),
        BLOCK_P=BLOCK_P,
        num_warps=4,
    )
    return output_tensor

# ----------------------------
# Self-test
# ----------------------------

def run_tests():
    tests = [
        {"batch": 1, "channels": 16, "kernelsize": 4, "seed": 4242, "size": 32},
        {"batch": 2, "channels": 16, "kernelsize": 4, "seed": 5236, "size": 32},
        {"batch": 1, "channels": 32, "kernelsize": 4, "seed": 1001, "size": 64},
        {"batch": 2, "channels": 32, "kernelsize": 8, "seed": 5531, "size": 64},
        {"batch": 1, "channels": 64, "kernelsize": 8, "seed": 9173, "size": 128},
    ]
    torch.cuda.synchronize()
    for t in tests:
        x, w, y = generate_input(**t)
        y_out = kernel_function(x, w, y)
        check_implementation((x, w, y), y_out)
    print("PASS")

if __name__ == "__main__":
    run_tests()
    sys.exit(0)


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)


# Ensure deterministic cuBLAS.
import os
if os.environ.get("CUBLAS_WORKSPACE_CONFIG", "") not in (":4096:8", ":16:8"):
    os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"

scrolls · 197 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 512024.

+ import os
+ import sys
+ from typing import Tuple, TypeVar, TypedDict
+
import torch
+ import torch.nn.functional as F
import triton
import triton.language as tl
+ # ----------------------------
+ # Original reference utilities
+ # ----------------------------
- # Triton kernel: 2D cross-correlation (PyTorch conv2d semantics) with:
- # - no padding
- # - stride = 1
- # - out_channels == in_channels == channels
- # Accumulation is in fp32 for both fp32 and bf16 inputs.
- #
- # Performance fix:
- # - Parallelize the output height (OH) across the launch grid instead of looping over OH inside
- # a single program. The previous version looped over OH in-kernel and timed out on large cases.
- # - Keep a deterministic sequential accumulation across (ic, kh, kw) to improve numerical
- # agreement with PyTorch's conv2d under strict tolerances, while avoiding excessive overhead.
+ def make_match_reference(reference: callable, **kwargs):
+ def wrapped(data, output):
+ return match_reference(data, output, reference=reference, **kwargs)
+ return wrapped
+ def match_reference(data, output, reference, rtol=1e-3, atol=1e-3):
+ ref = reference(data)
+ ok = torch.allclose(output, ref, rtol=rtol, atol=atol)
+ if not ok:
+ max_abs = (output - ref).abs().max().item()
+ max_rel = ((output - ref).abs() / (ref.abs() + 1e-12)).max().item()
+ raise AssertionError(f"Mismatch: max_abs={max_abs} max_rel={max_rel}")
+ return True
- @triton.autotune(
- configs=[
- triton.Config({'BLOCK_OC': 32, 'BLOCK_OW': 64}, num_stages=3, num_warps=4),
- triton.Config({'BLOCK_OC': 64, 'BLOCK_OW': 32}, num_stages=3, num_warps=8),
- triton.Config({'BLOCK_OC': 64, 'BLOCK_OW': 64}, num_stages=4, num_warps=8),
- triton.Config({'BLOCK_OC': 32, 'BLOCK_OW': 128}, num_stages=3, num_warps=4),
- triton.Config({'BLOCK_OC': 64, 'BLOCK_OW': 128}, num_stages=4, num_warps=8),
- ],
- key=['B', 'C', 'H', 'W', 'KH', 'KW'],
- )
- @triton.jit
- def _conv2d_nopad_stride1_parallel_oh(
- x_ptr, w_ptr, y_ptr,
- B, C, H, W, KH, KW, OH, OW,
- stride_x_b, stride_x_c, stride_x_h, stride_x_w,
- stride_w_oc, stride_w_ic, stride_w_kh, stride_w_kw,
- stride_y_b, stride_y_c, stride_y_h, stride_y_w,
- BLOCK_OC: tl.constexpr,
- BLOCK_OW: tl.constexpr,
- ):
- # Grid mapping:
- # axis 0: combined (batch, oh) tiles
- # axis 1: tiles over output channels
- # axis 2: tiles over output width
- pid_boh = tl.program_id(axis=0)
- pid_oc = tl.program_id(axis=1)
- pid_ow = tl.program_id(axis=2)
+ class DeterministicContext:
+ def __init__(self):
+ self.allow_tf32 = None
+ self.deterministic = None
+ self.cublas = None
- # Recover (b, oh) from packed pid_boh
- # We launch axis0 = B * OH, so:
- b = pid_boh // OH
- oh = pid_boh % OH
+ def __enter__(self):
+ self.cublas = os.environ.get("CUBLAS_WORKSPACE_CONFIG", "")
+ self.allow_tf32 = torch.backends.cudnn.allow_tf32
+ self.deterministic = torch.backends.cudnn.deterministic
+ torch.backends.cudnn.allow_tf32 = False
+ torch.backends.cudnn.deterministic = True
+ torch.use_deterministic_algorithms(True)
+ return self
- # Compute offsets for oc and ow tiles
- oc_offsets = pid_oc * BLOCK_OC + tl.arange(0, BLOCK_OC) # [BLOCK_OC]
- ow_offsets = pid_ow * BLOCK_OW + tl.arange(0, BLOCK_OW) # [BLOCK_OW]
+ def __exit__(self, exc_type, exc_value, traceback):
+ torch.backends.cudnn.allow_tf32 = self.allow_tf32
+ torch.backends.cudnn.deterministic = self.deterministic
+ torch.use_deterministic_algorithms(False)
+ os.environ["CUBLAS_WORKSPACE_CONFIG"] = self.cublas
- # Masks for bounds
- oc_mask = oc_offsets < C
- ow_mask = ow_offsets < OW
+ input_t = TypeVar("input_t", bound=Tuple[torch.Tensor, torch.Tensor, torch.Tensor])
+ output_t = TypeVar("output_t", bound=torch.Tensor)
- # FP32 accumulator
- acc = tl.zeros((BLOCK_OC, BLOCK_OW), dtype=tl.float32)
+ class TestSpec(TypedDict):
+ size: int
+ kernelsize: int
+ channels: int
+ batch: int
+ seed: int
- # Reduction across input channels, KH, KW in deterministic order
- for ic in tl.range(0, C):
- for kh in tl.range(0, KH):
- in_h = oh + kh # valid: no padding, stride=1
- for kw in tl.range(0, KW):
- in_w = ow_offsets + kw # [BLOCK_OW]
- in_w_mask = (ow_mask & (in_w < W))
+ def ref_kernel(data: input_t) -> output_t:
+ with DeterministicContext():
+ input_tensor, kernel, output = data
+ return F.conv2d(input_tensor, kernel, stride=1, padding=0)
- # Load weights for this (ic, kh, kw) across BLOCK_OC output channels
- w_ptrs = (
- w_ptr
- + oc_offsets * stride_w_oc
- + ic * stride_w_ic
- + kh * stride_w_kh
- + kw * stride_w_kw
- )
- w_vals = tl.load(w_ptrs, mask=oc_mask, other=0.0).to(tl.float32) # [BLOCK_OC]
+ def generate_input(size: int, kernelsize: int, channels: int, batch: int, seed: int) -> input_t:
+ gen = torch.Generator(device="cuda")
+ gen.manual_seed(seed)
+ x = torch.randn(batch, channels, size, size, device="cuda", dtype=torch.float32, generator=gen).contiguous()
+ w = torch.randn(channels, channels, kernelsize, kernelsize, device="cuda", dtype=torch.float32, generator=gen).contiguous()
+ y = torch.empty(batch, channels, size - kernelsize + 1, size - kernelsize + 1, device="cuda", dtype=torch.float32)
+ return x, w, y
- # Load input row for this (ic, in_h) across BLOCK_OW positions
- x_ptrs = (
- x_ptr
- + b * stride_x_b
- + ic * stride_x_c
- + in_h * stride_x_h
- + in_w * stride_x_w
- )
- x_vals = tl.load(x_ptrs, mask=in_w_mask, other=0.0).to(tl.float32) # [BLOCK_OW]
+ check_implementation = make_match_reference(ref_kernel, rtol=1e-3, atol=1e-3)
- # Outer product accumulate
- acc += w_vals[:, None] * x_vals[None, :]
+ # ----------------------------
+ # Triton implementation
+ # ----------------------------
- # Store results to output
- y_ptrs = (
- y_ptr
- + b * stride_y_b
- + oc_offsets[:, None] * stride_y_c
- + oh * stride_y_h
- + ow_offsets[None, :] * stride_y_w
- )
- y_mask = oc_mask[:, None] & ow_mask[None, :]
- tl.store(y_ptrs, acc.to(y_ptr.dtype.element_ty), mask=y_mask)
+ @triton.jit
+ def _conv2d_nchw_fwd_kernel(
+ x_ptr, w_ptr, y_ptr,
+ N: tl.constexpr, C: tl.constexpr, H: tl.constexpr, W: tl.constexpr,
+ KH: tl.constexpr, KW: tl.constexpr,
+ OH: tl.constexpr, OW: tl.constexpr,
+ stride_xn: tl.constexpr, stride_xc: tl.constexpr, stride_xh: tl.constexpr, stride_xw: tl.constexpr,
+ stride_wn: tl.constexpr, stride_wc: tl.constexpr, stride_wh: tl.constexpr, stride_ww: tl.constexpr,
+ stride_yn: tl.constexpr, stride_yc: tl.constexpr, stride_yh: tl.constexpr, stride_yw: tl.constexpr,
+ BLOCK_P: tl.constexpr, # pixels (OH*OW) per program
+ ):
+ pid_n = tl.program_id(0)
+ pid_co = tl.program_id(1)
+ pid_p = tl.program_id(2)
+ p0 = pid_p * BLOCK_P + tl.arange(0, BLOCK_P)
+ mask_p = p0 < (OH * OW)
- def kernel_function(input_tensor: torch.Tensor, weight: torch.Tensor, output_tensor: torch.Tensor = None):
- """
- Triton-backed 2D convolution (no padding, stride=1), matching PyTorch's F.conv2d semantics
- for the test settings:
- - input: [B, C, H, W]
- - weight: [C, C, KH, KW] (out_channels == in_channels == C)
- - output: [B, C, H-KH+1, W-KW+1]
+ oh = p0 // OW
+ ow = p0 - oh * OW
- All math runs inside the Triton kernel; the wrapper only validates/allocates/launches.
- Accumulation is done in fp32 for both fp32 and bf16 inputs.
+ acc = tl.zeros((BLOCK_P,), dtype=tl.float32)
- Args:
- input_tensor: CUDA tensor [B, C, H, W], dtype float32 or bfloat16, contiguous.
- weight: CUDA tensor [C, C, KH, KW], same dtype as input, contiguous.
- output_tensor (optional): Pre-allocated output [B, C, OH, OW], same dtype/device.
+ # Loop over C and KH*KW
+ for ci in range(0, C):
+ # base pointer for x at this (n, ci)
+ x_base_nc = x_ptr + pid_n * stride_xn + ci * stride_xc
+ # base pointer for w at this (co, ci)
+ w_base_coci = w_ptr + pid_co * stride_wn + ci * stride_wc
- Returns:
- Output tensor [B, C, OH, OW], CUDA, same dtype as input.
- """
- # Validate inputs
- assert input_tensor.is_cuda and weight.is_cuda, "Tensors must be CUDA"
- assert input_tensor.dtype in (torch.float32, torch.bfloat16), "Supported dtypes: float32, bfloat16"
- assert input_tensor.dtype == weight.dtype, "Input/weight dtypes must match"
- assert input_tensor.dim() == 4 and weight.dim() == 4, "Expected 4D tensors"
- assert input_tensor.is_contiguous() and weight.is_contiguous(), "Inputs must be contiguous"
+ for kh in range(0, KH):
+ for kw in range(0, KW):
+ x_ptrs = x_base_nc + (oh + kh) * stride_xh + (ow + kw) * stride_xw
+ w_val = tl.load(w_base_coci + kh * stride_wh + kw * stride_ww)
+ x_val = tl.load(x_ptrs, mask=mask_p, other=0.0)
+ acc += x_val * w_val
- B, C_in, H, W = input_tensor.shape
- OC, IC, KH, KW = weight.shape
- assert C_in == IC == OC, "Expect out_channels == in_channels == channels"
- assert KH >= 1 and KW >= 1, "Kernel dims must be positive"
+ y_ptrs = y_ptr + pid_n * stride_yn + pid_co * stride_yc + oh * stride_yh + ow * stride_yw
+ tl.store(y_ptrs, acc, mask=mask_p)
+ def kernel_function(input_tensor: torch.Tensor, kernel: torch.Tensor, output_tensor: torch.Tensor) -> torch.Tensor:
+ # Validate (no torch math ops)
+ if input_tensor.device.type != "cuda" or kernel.device.type != "cuda" or output_tensor.device.type != "cuda":
+ raise ValueError("All tensors must be on CUDA")
+ if input_tensor.dtype != torch.float32 or kernel.dtype != torch.float32 or output_tensor.dtype != torch.float32:
+ raise ValueError("Only float32 supported")
+ if not input_tensor.is_contiguous() or not kernel.is_contiguous() or not output_tensor.is_contiguous():
+ raise ValueError("All tensors must be contiguous")
+ if input_tensor.ndim != 4 or kernel.ndim != 4 or output_tensor.ndim != 4:
+ raise ValueError("Expected NCHW tensors")
+ N, C, H, W = input_tensor.shape
+ Cout, Cin, KH, KW = kernel.shape
+ if Cout != C or Cin != C:
+ raise ValueError("This problem expects kernel shape [C, C, KH, KW]")
OH = H - KH + 1
OW = W - KW + 1
- assert OH > 0 and OW > 0, "Invalid kernel size: no padding, stride=1"
+ if output_tensor.shape != (N, Cout, OH, OW):
+ raise ValueError("output_tensor has incorrect shape")
- if output_tensor is None:
- output_tensor = torch.empty((B, OC, OH, OW), device=input_tensor.device, dtype=input_tensor.dtype)
- else:
- assert output_tensor.is_cuda, "Output tensor must be CUDA"
- assert output_tensor.dtype == input_tensor.dtype, "Output dtype mismatch"
- assert output_tensor.is_contiguous(), "Output must be contiguous"
- assert tuple(output_tensor.shape) == (B, OC, OH, OW), "Output shape mismatch"
+ # Launch: 3D grid (N, Cout, blocks over OH*OW)
+ BLOCK_P = 256
- # Strides
- stride_x_b, stride_x_c, stride_x_h, stride_x_w = input_tensor.stride()
- stride_w_oc, stride_w_ic, stride_w_kh, stride_w_kw = weight.stride()
- stride_y_b, stride_y_c, stride_y_h, stride_y_w = output_tensor.stride()
+ def grid(meta):
+ return (N, Cout, triton.cdiv(OH * OW, meta["BLOCK_P"]))
- # 3D grid over (batch*OH, channel tiles, output-width tiles)
- def grid(META):
- return (
- B * OH,
- triton.cdiv(C_in, META['BLOCK_OC']),
- triton.cdiv(OW, META['BLOCK_OW']),
- )
-
- # Launch Triton kernel
- _conv2d_nopad_stride1_parallel_oh[grid](
- input_tensor, weight, output_tensor,
- B, C_in, H, W, KH, KW, OH, OW,
- stride_x_b, stride_x_c, stride_x_h, stride_x_w,
- stride_w_oc, stride_w_ic, stride_w_kh, stride_w_kw,
- stride_y_b, stride_y_c, stride_y_h, stride_y_w,
+ _conv2d_nchw_fwd_kernel[grid](
+ input_tensor, kernel, output_tensor,
+ N=N, C=C, H=H, W=W, KH=KH, KW=KW, OH=OH, OW=OW,
+ stride_xn=input_tensor.stride(0), stride_xc=input_tensor.stride(1),
+ stride_xh=input_tensor.stride(2), stride_xw=input_tensor.stride(3),
+ stride_wn=kernel.stride(0), stride_wc=kernel.stride(1),
+ stride_wh=kernel.stride(2), stride_ww=kernel.stride(3),
+ stride_yn=output_tensor.stride(0), stride_yc=output_tensor.stride(1),
+ stride_yh=output_tensor.stride(2), stride_yw=output_tensor.stride(3),
+ BLOCK_P=BLOCK_P,
+ num_warps=4,
)
-
return output_tensor
+ # ----------------------------
+ # Self-test
+ # ----------------------------
+
+ def run_tests():
+ tests = [
+ {"batch": 1, "channels": 16, "kernelsize": 4, "seed": 4242, "size": 32},
+ {"batch": 2, "channels": 16, "kernelsize": 4, "seed": 5236, "size": 32},
+ {"batch": 1, "channels": 32, "kernelsize": 4, "seed": 1001, "size": 64},
+ {"batch": 2, "channels": 32, "kernelsize": 8, "seed": 5531, "size": 64},
+ {"batch": 1, "channels": 64, "kernelsize": 8, "seed": 9173, "size": 128},
+ ]
+ torch.cuda.synchronize()
+ for t in tests:
+ x, w, y = generate_input(**t)
+ y_out = kernel_function(x, w, y)
+ check_implementation((x, w, y), y_out)
+ print("PASS")
+
+ if __name__ == "__main__":
+ run_tests()
+ sys.exit(0)
+
+
import inspect
def custom_kernel(input):
scrolls · 324 diff lines total

Best evidence level for this revision: reported

JSON