COURSE / SOURCE

check_against_pytorch.py

All lessons
Source filecode/day83-cudnn/check_against_pytorch.py

This is the source used by the lesson and its recorded evidence. Compile commands and expected output live in the directory README.

# SPDX-License-Identifier: MIT
#
# Day 83: the PyTorch half of the Conv2D check.
#
# conv_graph.cu ran one convolution through the cuDNN graph API and wrote its
# input, filter and output to disk. This runs the same convolution through
# torch.nn.functional.conv2d and compares the two answers under the tolerance
# the day states, which is the accumulation-depth bound from
# research/EXERCISE-DESIGN.md and not a number picked to make the test pass.
#
# Both sides use cuDNN's convolution semantics, but they need not load the
# same build: the captured standalone run used cuDNN 9.13.1 and the torch
# wheel bundled 9.10.2. A disagreement still points first at our graph,
# layout or convolution mode, while an exact match is stronger than assuming
# the two processes shared one implementation.
#
# Pinned: torch 2.13.0+cu126 (cp312 manylinux wheel, the newest cu126 build in
# https://download.pytorch.org/whl/cu126/torch/ , checked 2026-09-01). numpy is
# whatever that wheel resolves; record `pip freeze | grep -Ei "torch|numpy"` in
# the evidence file, because the torch wheel bundles its own cuDNN and its
# version is part of the result.
#
# Run: python3 check_against_pytorch.py
#
# VERIFIED: torch 2.13.0+cu126 with bundled cuDNN 9.10.2 on a Tesla T4,
# 2026-09-02. See evidence/torch-2026-09-02.txt.

import math
import sys

import numpy as np
import torch


def read_shapes(path):
    """The .cu wrote these; repeating them here would let the two drift."""
    shapes = {}
    with open(path, encoding="utf-8") as f:
        for line in f:
            key, value = line.split()
            shapes[key] = int(value)
    return shapes


def main():
    s = read_shapes("conv_shapes.txt")

    print(f"torch {torch.__version__}, built against CUDA "
          f"{torch.version.cuda}, cuDNN {torch.backends.cudnn.version()}")
    if not torch.cuda.is_available():
        print("no CUDA device visible to torch", file=sys.stderr)
        return 1
    print(f"GPU: {torch.cuda.get_device_name(0)}")

    # Pin the math mode. TF32 is on by default for cuDNN convolutions on
    # Ampere and newer, and it would silently drop the inputs to 10 mantissa
    # bits, which is a different computation from the fp32 one conv_graph.cu
    # asked cuDNN for. A T4 has no TF32 path so this changes nothing there and
    # everything on an A100.
    torch.backends.cudnn.allow_tf32 = False
    torch.backends.cudnn.benchmark = False
    print(f"cudnn.allow_tf32 = {torch.backends.cudnn.allow_tf32}, "
          f"dtype float32, accumulate float32")

    x = np.fromfile("conv_x.f32", dtype=np.float32).reshape(
        s["n"], s["c"], s["h"], s["w"])
    w = np.fromfile("conv_w.f32", dtype=np.float32).reshape(
        s["k"], s["c"], s["r"], s["s"])
    ours = np.fromfile("conv_y.f32", dtype=np.float32).reshape(
        s["n"], s["k"], s["p"], s["q"])

    # F.conv2d is cross-correlation, the same thing CUDNN_CROSS_CORRELATION
    # names. torch has no flag for the flipped kind at all.
    xt = torch.from_numpy(x).cuda()
    wt = torch.from_numpy(w).cuda()
    yt = torch.nn.functional.conv2d(
        xt, wt, bias=None, stride=s["stride"], padding=s["pad"],
        dilation=s["dilation"])
    theirs = yt.cpu().numpy()

    if theirs.shape != ours.shape:
        print(f"shape mismatch: cuDNN {ours.shape}, torch {theirs.shape}",
              file=sys.stderr)
        return 1

    # Same bound as the .cu, for the same reason: the two runs accumulate
    # c*r*s products and are free to do it in different orders.
    k_terms = s["c"] * s["r"] * s["s"]
    eps = 2.0 ** -23
    rtol = max(1e-5, 4.0 * eps * math.sqrt(k_terms))
    atol = rtol * float(np.abs(theirs).max())
    print(f"tolerance |cuDNN-torch| <= {atol:.3e} + {rtol:.3e}*|torch|")
    print(f"          rtol = max(1e-5, 4*2^-23*sqrt({k_terms}))")

    diff = np.abs(ours.astype(np.float64) - theirs.astype(np.float64))
    bound = atol + rtol * np.abs(theirs.astype(np.float64))
    bad = int(np.count_nonzero(diff > bound))
    flat = int(np.argmax(diff))
    print(f"largest |cuDNN - torch| = {diff.max():.3e} at "
          f"{np.unravel_index(flat, diff.shape)}")

    if bad:
        first = int(np.argmax(diff > bound))
        idx = np.unravel_index(first, diff.shape)
        print(f"first mismatch at {idx}: cuDNN {ours[idx]:.9g}, "
              f"torch {theirs[idx]:.9g}", file=sys.stderr)
        print(f"{bad} of {diff.size} outputs outside tolerance",
              file=sys.stderr)
        return 1

    print(f"all {diff.size} outputs inside tolerance")
    return 0


if __name__ == "__main__":
    sys.exit(main())