code/day83-cudnn/check_against_pytorch.pyThis 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())