Skip to content
KernelIndex
Search⌘K

submission 550064

KernelAgent · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

gpumode_submit_1kb2ipfs.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-conv2d-v2-550064?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
136.1ms
#22 of 35
2026-03-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8e91485d2d64daca70aec431872bb7831257fff366484d562d2c924d4b10db27
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, num_stages=2,
stages = 2num_warps=4, num_stages=2,

Kernel source

gpumode_submit_1kb2ipfs.py189 lines
import torch
import triton
import triton.language as tl


@triton.jit
def _conv2d_nchw_str1_nopad_kernel(
    x_ptr,  # *const T_x
    w_ptr,  # *const T_w
    y_ptr,  # *T_y
    # sizes
    N, C_in, H, W,
    C_out, K, H_out, W_out,
    # strides for input x (NCHW)
    stride_xn, stride_xc, stride_xh, stride_xw,
    # strides for weight w (O, I, KH, KW)
    stride_wo, stride_wi, stride_wkh, stride_wkw,
    # strides for output y (N, O, H_out, W_out)
    stride_yn, stride_yc, stride_yh, stride_yw,
    # compile-time tile/block params
    BLOCK_H: tl.constexpr,
    BLOCK_W: tl.constexpr,
    BLOCK_OC: tl.constexpr,
    M: tl.constexpr,  # = BLOCK_H * BLOCK_W (compile-time)
):
    # Program IDs
    pid_n = tl.program_id(0)
    pid_sp = tl.program_id(1)
    pid_ocg = tl.program_id(2)

    # Tile decomposition for spatial output grid
    num_tiles_w = tl.cdiv(W_out, BLOCK_W)
    tile_h_idx = pid_sp // num_tiles_w
    tile_w_idx = pid_sp % num_tiles_w
    oh0 = tile_h_idx * BLOCK_H
    ow0 = tile_w_idx * BLOCK_W

    # Output channel group origin
    oc0 = pid_ocg * BLOCK_OC

    # Output channel range and mask
    oc_range = oc0 + tl.arange(0, BLOCK_OC)
    oc_mask = oc_range < C_out

    # Flatten spatial tile indexing: indices 0..M-1 map to (mh, mw)
    m_idx = tl.arange(0, M)
    mh = m_idx // BLOCK_W  # [0..BLOCK_H-1]
    mw = m_idx % BLOCK_W   # [0..BLOCK_W-1]

    # Per-lane output coords and mask for in-bounds (avoid indexing tensors with tensors)
    oh_m = oh0 + mh
    ow_m = ow0 + mw
    mask_m = (oh_m < H_out) & (ow_m < W_out)

    # Base pointers for this program instance
    x_n_base = x_ptr + pid_n * stride_xn
    y_n_base = y_ptr + pid_n * stride_yn

    # Accumulator (M x BLOCK_OC) in fp32 for numerical stability
    acc = tl.zeros((M, BLOCK_OC), dtype=tl.float32)

    # Convolution:
    # y[n, oc, oh, ow] = sum_{ic, kh, kw} x[n, ic, oh+kh, ow+kw] * w[oc, ic, kh, kw]
    for ic in tl.range(0, C_in):
        x_ic_base = x_n_base + ic * stride_xc
        w_ic_base = w_ptr + ic * stride_wi
        for kh in tl.range(0, K):
            ih_m = oh_m + kh  # (M,)
            x_ptrs_row = x_ic_base + ih_m * stride_xh  # (M,)
            for kw in tl.range(0, K):
                iw_m = ow_m + kw  # (M,)
                x_ptrs_m = x_ptrs_row + iw_m * stride_xw  # (M,)
                x_vec = tl.load(x_ptrs_m, mask=mask_m, other=0.0)

                # Weight vector for this (ic, kh, kw) and output channel group
                w_ptrs_vec = (
                    w_ic_base
                    + oc0 * stride_wo
                    + kh * stride_wkh
                    + kw * stride_wkw
                    + tl.arange(0, BLOCK_OC) * stride_wo
                )
                w_vec = tl.load(w_ptrs_vec, mask=oc_mask, other=0.0)

                # Outer product accumulation into (M, BLOCK_OC)
                acc += x_vec.to(tl.float32)[:, None] * w_vec.to(tl.float32)[None, :]

    # Store results
    y_ptrs = (
        y_n_base
        + oc_range[None, :] * stride_yc
        + oh_m[:, None] * stride_yh
        + ow_m[:, None] * stride_yw
    )
    out_dtype = y_ptr.dtype.element_ty
    y_vals = acc.to(out_dtype)
    y_mask = mask_m[:, None] & oc_mask[None, :]
    tl.store(y_ptrs, y_vals, mask=y_mask)


def kernel_function(input_tensor: torch.Tensor, kernel: torch.Tensor, out: torch.Tensor = None):
    """
    Triton 2D convolution (NCHW) with stride=1 and padding=0.
    - Single fused kernel: performs full accumulation over input channels and kernel spatial dims.
    - No math or reductions in the wrapper: only validation, allocation, and launch.

    Args:
        input_tensor: (N, C_in, H, W), CUDA tensor (float32 / float16 / bfloat16)
        kernel:       (C_out, C_in, K, K), contiguous CUDA tensor, same dtype/device as input
        out:          Optional preallocated tensor (N, C_out, H-K+1, W-K+1)

    Returns:
        (N, C_out, H-K+1, W-K+1) tensor on CUDA with same dtype as input.
    """
    # Validation (no compute)
    assert isinstance(input_tensor, torch.Tensor) and isinstance(kernel, torch.Tensor)
    assert input_tensor.is_cuda and kernel.is_cuda
    assert input_tensor.dim() == 4 and kernel.dim() == 4
    assert input_tensor.device == kernel.device
    assert input_tensor.dtype == kernel.dtype
    assert input_tensor.is_contiguous() and kernel.is_contiguous()

    N, C_in, H, W = input_tensor.shape
    C_out, C_in_w, K, K_w = kernel.shape
    assert C_in == C_in_w and K == K_w, "Incompatible kernel shape"
    assert H >= K and W >= K, "Kernel larger than input"
    H_out = H - K + 1
    W_out = W - K + 1

    # Supported dtypes
    if input_tensor.dtype not in (torch.float32, torch.float16, torch.bfloat16):
        raise TypeError(f"Unsupported dtype: {input_tensor.dtype}")
    out_dtype = input_tensor.dtype

    # Allocate output if needed
    if out is None:
        out = torch.empty((N, C_out, H_out, W_out), device=input_tensor.device, dtype=out_dtype)
    else:
        assert isinstance(out, torch.Tensor) and out.is_cuda
        assert out.device == input_tensor.device
        assert out.dtype == out_dtype
        assert out.shape == (N, C_out, H_out, W_out)
        assert out.is_contiguous()

    # Element-wise strides
    stride_xn, stride_xc, stride_xh, stride_xw = input_tensor.stride()
    stride_wo, stride_wi, stride_wkh, stride_wkw = kernel.stride()
    stride_yn, stride_yc, stride_yh, stride_yw = out.stride()

    # Tile sizes (powers of two, masks handle tails)
    BLOCK_H = 16
    BLOCK_W = 16
    BLOCK_OC = 32
    M = BLOCK_H * BLOCK_W  # constexpr

    def grid(meta):
        num_tiles_h = triton.cdiv(H_out, meta["BLOCK_H"])
        num_tiles_w = triton.cdiv(W_out, meta["BLOCK_W"])
        return (
            N,                                       # batch
            num_tiles_h * num_tiles_w,               # spatial tiles
            triton.cdiv(C_out, meta["BLOCK_OC"]),    # output channel tiles
        )

    # Launch kernel
    _conv2d_nchw_str1_nopad_kernel[grid](
        input_tensor, kernel, out,
        N, C_in, H, W,
        C_out, K, H_out, W_out,
        stride_xn, stride_xc, stride_xh, stride_xw,
        stride_wo, stride_wi, stride_wkh, stride_wkw,
        stride_yn, stride_yc, stride_yh, stride_yw,
        BLOCK_H=BLOCK_H, BLOCK_W=BLOCK_W, BLOCK_OC=BLOCK_OC, M=M,
        num_warps=4, num_stages=2,
    )
    return out

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 · 189 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 513456.

- 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
+ @triton.jit
+ def _conv2d_nchw_str1_nopad_kernel(
+ x_ptr, # *const T_x
+ w_ptr, # *const T_w
+ y_ptr, # *T_y
+ # sizes
+ N, C_in, H, W,
+ C_out, K, H_out, W_out,
+ # strides for input x (NCHW)
+ stride_xn, stride_xc, stride_xh, stride_xw,
+ # strides for weight w (O, I, KH, KW)
+ stride_wo, stride_wi, stride_wkh, stride_wkw,
+ # strides for output y (N, O, H_out, W_out)
+ stride_yn, stride_yc, stride_yh, stride_yw,
+ # compile-time tile/block params
+ BLOCK_H: tl.constexpr,
+ BLOCK_W: tl.constexpr,
+ BLOCK_OC: tl.constexpr,
+ M: tl.constexpr, # = BLOCK_H * BLOCK_W (compile-time)
+ ):
+ # Program IDs
+ pid_n = tl.program_id(0)
+ pid_sp = tl.program_id(1)
+ pid_ocg = tl.program_id(2)
- 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
+ # Tile decomposition for spatial output grid
+ num_tiles_w = tl.cdiv(W_out, BLOCK_W)
+ tile_h_idx = pid_sp // num_tiles_w
+ tile_w_idx = pid_sp % num_tiles_w
+ oh0 = tile_h_idx * BLOCK_H
+ ow0 = tile_w_idx * BLOCK_W
- class DeterministicContext:
- def __init__(self):
- self.allow_tf32 = None
- self.deterministic = None
- self.cublas = None
+ # Output channel group origin
+ oc0 = pid_ocg * BLOCK_OC
- 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
+ # Output channel range and mask
+ oc_range = oc0 + tl.arange(0, BLOCK_OC)
+ oc_mask = oc_range < C_out
- 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
+ # Flatten spatial tile indexing: indices 0..M-1 map to (mh, mw)
+ m_idx = tl.arange(0, M)
+ mh = m_idx // BLOCK_W # [0..BLOCK_H-1]
+ mw = m_idx % BLOCK_W # [0..BLOCK_W-1]
- input_t = TypeVar("input_t", bound=Tuple[torch.Tensor, torch.Tensor, torch.Tensor])
- output_t = TypeVar("output_t", bound=torch.Tensor)
+ # Per-lane output coords and mask for in-bounds (avoid indexing tensors with tensors)
+ oh_m = oh0 + mh
+ ow_m = ow0 + mw
+ mask_m = (oh_m < H_out) & (ow_m < W_out)
- class TestSpec(TypedDict):
- size: int
- kernelsize: int
- channels: int
- batch: int
- seed: int
+ # Base pointers for this program instance
+ x_n_base = x_ptr + pid_n * stride_xn
+ y_n_base = y_ptr + pid_n * stride_yn
- 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)
+ # Accumulator (M x BLOCK_OC) in fp32 for numerical stability
+ acc = tl.zeros((M, BLOCK_OC), dtype=tl.float32)
- 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
+ # Convolution:
+ # y[n, oc, oh, ow] = sum_{ic, kh, kw} x[n, ic, oh+kh, ow+kw] * w[oc, ic, kh, kw]
+ for ic in tl.range(0, C_in):
+ x_ic_base = x_n_base + ic * stride_xc
+ w_ic_base = w_ptr + ic * stride_wi
+ for kh in tl.range(0, K):
+ ih_m = oh_m + kh # (M,)
+ x_ptrs_row = x_ic_base + ih_m * stride_xh # (M,)
+ for kw in tl.range(0, K):
+ iw_m = ow_m + kw # (M,)
+ x_ptrs_m = x_ptrs_row + iw_m * stride_xw # (M,)
+ x_vec = tl.load(x_ptrs_m, mask=mask_m, other=0.0)
- check_implementation = make_match_reference(ref_kernel, rtol=1e-3, atol=1e-3)
+ # Weight vector for this (ic, kh, kw) and output channel group
+ w_ptrs_vec = (
+ w_ic_base
+ + oc0 * stride_wo
+ + kh * stride_wkh
+ + kw * stride_wkw
+ + tl.arange(0, BLOCK_OC) * stride_wo
+ )
+ w_vec = tl.load(w_ptrs_vec, mask=oc_mask, other=0.0)
- # ----------------------------
- # Triton implementation
- # ----------------------------
+ # Outer product accumulation into (M, BLOCK_OC)
+ acc += x_vec.to(tl.float32)[:, None] * w_vec.to(tl.float32)[None, :]
- @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)
+ # Store results
+ y_ptrs = (
+ y_n_base
+ + oc_range[None, :] * stride_yc
+ + oh_m[:, None] * stride_yh
+ + ow_m[:, None] * stride_yw
+ )
+ out_dtype = y_ptr.dtype.element_ty
+ y_vals = acc.to(out_dtype)
+ y_mask = mask_m[:, None] & oc_mask[None, :]
+ tl.store(y_ptrs, y_vals, mask=y_mask)
- p0 = pid_p * BLOCK_P + tl.arange(0, BLOCK_P)
- mask_p = p0 < (OH * OW)
- oh = p0 // OW
- ow = p0 - oh * OW
+ def kernel_function(input_tensor: torch.Tensor, kernel: torch.Tensor, out: torch.Tensor = None):
+ """
+ Triton 2D convolution (NCHW) with stride=1 and padding=0.
+ - Single fused kernel: performs full accumulation over input channels and kernel spatial dims.
+ - No math or reductions in the wrapper: only validation, allocation, and launch.
- acc = tl.zeros((BLOCK_P,), dtype=tl.float32)
+ Args:
+ input_tensor: (N, C_in, H, W), CUDA tensor (float32 / float16 / bfloat16)
+ kernel: (C_out, C_in, K, K), contiguous CUDA tensor, same dtype/device as input
+ out: Optional preallocated tensor (N, C_out, H-K+1, W-K+1)
- # 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:
+ (N, C_out, H-K+1, W-K+1) tensor on CUDA with same dtype as input.
+ """
+ # Validation (no compute)
+ assert isinstance(input_tensor, torch.Tensor) and isinstance(kernel, torch.Tensor)
+ assert input_tensor.is_cuda and kernel.is_cuda
+ assert input_tensor.dim() == 4 and kernel.dim() == 4
+ assert input_tensor.device == kernel.device
+ assert input_tensor.dtype == kernel.dtype
+ assert input_tensor.is_contiguous() and kernel.is_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
+ N, C_in, H, W = input_tensor.shape
+ C_out, C_in_w, K, K_w = kernel.shape
+ assert C_in == C_in_w and K == K_w, "Incompatible kernel shape"
+ assert H >= K and W >= K, "Kernel larger than input"
+ H_out = H - K + 1
+ W_out = W - K + 1
- 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)
+ # Supported dtypes
+ if input_tensor.dtype not in (torch.float32, torch.float16, torch.bfloat16):
+ raise TypeError(f"Unsupported dtype: {input_tensor.dtype}")
+ out_dtype = input_tensor.dtype
- 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")
+ # Allocate output if needed
+ if out is None:
+ out = torch.empty((N, C_out, H_out, W_out), device=input_tensor.device, dtype=out_dtype)
+ else:
+ assert isinstance(out, torch.Tensor) and out.is_cuda
+ assert out.device == input_tensor.device
+ assert out.dtype == out_dtype
+ assert out.shape == (N, C_out, H_out, W_out)
+ assert out.is_contiguous()
- # Launch: 3D grid (N, Cout, blocks over OH*OW)
- BLOCK_P = 256
+ # Element-wise strides
+ stride_xn, stride_xc, stride_xh, stride_xw = input_tensor.stride()
+ stride_wo, stride_wi, stride_wkh, stride_wkw = kernel.stride()
+ stride_yn, stride_yc, stride_yh, stride_yw = out.stride()
+ # Tile sizes (powers of two, masks handle tails)
+ BLOCK_H = 16
+ BLOCK_W = 16
+ BLOCK_OC = 32
+ M = BLOCK_H * BLOCK_W # constexpr
+
def grid(meta):
- return (N, Cout, triton.cdiv(OH * OW, meta["BLOCK_P"]))
+ num_tiles_h = triton.cdiv(H_out, meta["BLOCK_H"])
+ num_tiles_w = triton.cdiv(W_out, meta["BLOCK_W"])
+ return (
+ N, # batch
+ num_tiles_h * num_tiles_w, # spatial tiles
+ triton.cdiv(C_out, meta["BLOCK_OC"]), # output channel tiles
+ )
- _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,
+ # Launch kernel
+ _conv2d_nchw_str1_nopad_kernel[grid](
+ input_tensor, kernel, out,
+ N, C_in, H, W,
+ C_out, K, H_out, W_out,
+ stride_xn, stride_xc, stride_xh, stride_xw,
+ stride_wo, stride_wi, stride_wkh, stride_wkw,
+ stride_yn, stride_yc, stride_yh, stride_yw,
+ BLOCK_H=BLOCK_H, BLOCK_W=BLOCK_W, BLOCK_OC=BLOCK_OC, M=M,
+ num_warps=4, num_stages=2,
)
- return output_tensor
+ return out
- # ----------------------------
- # 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 · 344 diff lines total

Best evidence level for this revision: reported

JSON