code/day88-torch-profiler/profile_model.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 88: profiling a small PyTorch model down to the kernel.
#
# One block, two spellings. A Linear (a GEMM) followed by day 48's
# elementwise chain, written first the way PyTorch spells it and then as the
# single fused operator in fused_chain.cu. The program checks that the two
# agree, times both with CUDA events, profiles both with torch.profiler, and
# wraps every phase in an NVTX range so the same run is readable in Nsight
# Systems, which is day 41's lesson arriving in the ML world.
#
# Requirements, and why each pin is what it is (README.md has the commands):
# torch==2.13.0 from the cu126 index. The Tesla T4 is compute capability
# 7.5; PyTorch's CUDA support matrix lists Turing(7.5) for the 12.6, 13.0
# and 13.2 builds, so any of the three would run, and cu126 is the one
# whose major matches the verification node's nvcc 12.6. That match is not
# cosmetic: cpp_extension calls that nvcc to build fused_chain.cu.
# ninja, because cpp_extension shells out to it.
#
# Run: python3 profile_model.py
#
# VERIFIED: Tesla T4, driver 580.173.02, CUDA 12.6 and
# torch 2.13.0+cu126, 2026-09-02.
import os
import sys
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.profiler import ProfilerActivity, profile
from torch.utils.cpp_extension import load
HERE = os.path.dirname(os.path.abspath(__file__))
PROFILE_DIR = os.path.join(HERE, "profile")
# 8192 x 1024 floats is 32 MiB per tensor, which is eight times the T4's 4
# MiB L2, so the elementwise chain has to go to DRAM on every stage and the
# byte model below is not describing cache traffic. The GEMM at this shape is
# 2 * 8192 * 1024 * 1024 flops, close enough to the chain's cost that neither
# one trivially owns the profile.
BATCH, DIM = 8192, 1024
WARMUPS, RUNS = 3, 10
# day 48's constants, matching fused_chain.cu. Both files have to change
# together or the correctness check below is comparing two functions.
GAMMA, BETA = 1.5, -0.25
LO, HI = -6.0, 3.0
# Floats moved per element, counted off the ops themselves.
# eager: mul 2, add-bias 2, gelu 2, add-residual 3 (it reads two
# tensors), clamp 2 -> 11
# fused: x, the residual, the output -> 3
# Day 48's hand-staged chain moved 9, not 11, because it did the scale and
# the bias in one kernel. Eager PyTorch cannot: `h * GAMMA + BETA` is two
# operators, so it is two launches and one extra round trip.
EAGER_FLOATS_PER_ELEM = 11
FUSED_FLOATS_PER_ELEM = 3
# The two chains run the same float32 operations in the same order, so the
# only thing that can move a result is the compiler contracting `v * GAMMA +
# BETA` into one FMA inside the kernel while eager rounds twice. That is one
# rounding of a float32 intermediate, about 2^-24 relative, and GELU's
# derivative is under 1.2 everywhere, so nothing amplifies it. There is no
# reduction here, so the bound does not grow with BATCH * DIM. Day 48
# measured its own staged-against-fused pair at 2.38e-07 absolute worst
# case; these are day 48's tolerances, two orders above that.
RTOL, ATOL = 1e-5, 1e-6
def banner():
"""Print what is about to be measured, from the machine, not from a
README that can go stale."""
print(f"torch {torch.__version__}, built against CUDA "
f"{torch.version.cuda}")
if not torch.cuda.is_available():
print("no CUDA device visible; this day needs one, exiting 1")
return False
props = torch.cuda.get_device_properties(0)
print(f"GPU: {props.name} (compute capability {props.major}."
f"{props.minor}), {props.multi_processor_count} SMs, "
f"{props.total_memory // (1024 * 1024)} MiB")
print(f"driver-visible CUDA_HOME: {os.environ.get('CUDA_HOME', 'unset')}")
print(f"shape: {BATCH} x {DIM} = {BATCH * DIM} elements, "
f"{BATCH * DIM * 4 // (1024 * 1024)} MiB per float32 tensor")
return True
# snippet: eager-chain
class Block(nn.Module):
"""A Linear, then day 48's chain. `fused=True` swaps the five
elementwise operators for one call into fused_chain.cu."""
def __init__(self, dim, fused):
super().__init__()
self.fc = nn.Linear(dim, dim, bias=True)
self.fused = fused
def forward(self, x, residual):
h = self.fc(x)
if self.fused:
return torch.ops.day88.fused_chain(h, residual)
h = h * GAMMA + BETA
h = F.gelu(h, approximate="tanh")
h = h + residual
return h.clamp(LO, HI)
# end snippet: eager-chain
def build_extension():
"""Compile fused_chain.cu and register torch.ops.day88.fused_chain.
is_python_module=False because the file binds through TORCH_LIBRARY
rather than pybind11, so there is no module object to import: the
operator arrives on the torch.ops namespace instead. That is also what
keeps it visible to torch.compile, which day 87 covers.
"""
load(
name="day88_fused_chain",
sources=[os.path.join(HERE, "fused_chain.cu")],
extra_cuda_cflags=["-O3", "-lineinfo"],
is_python_module=False,
verbose=True,
)
def correctness(eager, fused, x, residual):
"""Compare the two blocks element by element. Returns True on pass and
prints the first disagreement on failure."""
with torch.no_grad():
want = eager(x, residual)
got = fused(x, residual)
diff = (got - want).abs()
tol = RTOL * want.abs() + ATOL
bad = diff > tol
count = int(bad.sum().item())
worst = float(diff.max().item())
clamped = int(((want <= LO) | (want >= HI)).sum().item())
print(f"correctness: {count} elements outside rtol {RTOL} atol {ATOL}, "
f"worst absolute difference {worst:.3e}")
print(f" {clamped} of {want.numel()} outputs sit on the "
f"clamp, so that stage is doing something")
if count:
idx = tuple(bad.nonzero()[0].tolist())
print(f" first at {idx}: eager {want[idx].item()!r} "
f"fused {got[idx].item()!r}")
return False
return True
def time_block(block, x, residual):
"""Mean milliseconds over RUNS timed runs after WARMUPS, on CUDA events,
because a host clock around an asynchronous launch times the launch.
Day 9 is the lesson."""
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
with torch.no_grad():
for _ in range(WARMUPS):
block(x, residual)
torch.cuda.synchronize()
start.record()
for _ in range(RUNS):
block(x, residual)
end.record()
torch.cuda.synchronize()
return start.elapsed_time(end) / RUNS
def profile_block(block, x, residual, name):
"""One profiled forward pass, with the table printed two ways and the
Chrome trace written out.
record_shapes=True is what makes the second table possible: it keeps each
operator's input shapes, so `aten::add` called on two different shapes
stops being one row. That is how you find which of five identical-looking
elementwise calls is the expensive one.
"""
with torch.no_grad():
block(x, residual) # warm the allocator and the kernels
torch.cuda.synchronize()
with profile(activities=[ProfilerActivity.CPU,
ProfilerActivity.CUDA],
record_shapes=True) as prof:
block(x, residual)
torch.cuda.synchronize()
print(f"\n--- {name}: by self CUDA time ---")
print(prof.key_averages().table(sort_by="self_cuda_time_total",
row_limit=15))
print(f"\n--- {name}: same run, grouped by input shape ---")
print(prof.key_averages(group_by_input_shape=True).table(
sort_by="self_cuda_time_total", row_limit=10))
path = os.path.join(PROFILE_DIR, f"day88-{name}-trace.json")
prof.export_chrome_trace(path)
print(f"chrome trace: {path}")
# snippet: nvtx-pass
def nvtx_pass(eager, fused, x, residual):
"""The same two blocks again, under NVTX ranges and nothing else.
torch.profiler already knows the operator names. Nsight Systems does not:
without these ranges an nsys capture of a model is one long row of
anonymous cudaLaunchKernel bars, which is exactly the problem day 41
opened with. The other route is
torch.autograd.profiler.emit_nvtx(), which puts a range around every
autograd operation automatically; it is finer and noisier than this.
"""
with torch.no_grad():
for _ in range(RUNS):
with torch.cuda.nvtx.range("eager-block"):
eager(x, residual)
with torch.cuda.nvtx.range("fused-block"):
fused(x, residual)
torch.cuda.synchronize()
# end snippet: nvtx-pass
def main():
if not banner():
return 1
os.makedirs(PROFILE_DIR, exist_ok=True)
torch.manual_seed(88)
build_extension()
device = torch.device("cuda:0")
eager = Block(DIM, fused=False).to(device).eval()
fused = Block(DIM, fused=True).to(device).eval()
# One set of weights, two spellings of the same block. Without this the
# two paths would be computing different functions and the comparison
# below would mean nothing.
fused.load_state_dict(eager.state_dict())
x = torch.randn(BATCH, DIM, device=device)
# Uniform on [-2, 2) rather than normal, so the residual pushes a real
# fraction of the outputs onto the clamp instead of a handful of tails.
residual = torch.empty(BATCH, DIM, device=device).uniform_(-2.0, 2.0)
ok = correctness(eager, fused, x, residual)
eager_ms = time_block(eager, x, residual)
fused_ms = time_block(fused, x, residual)
print("\nmean of "
f"{RUNS} runs after {WARMUPS} warm-ups, CUDA events")
print(f" eager block {eager_ms:9.4f} ms")
print(f" fused block {fused_ms:9.4f} ms")
print(f" eager / fused {eager_ms / fused_ms:9.4f}")
print(f" chain floats per element, eager {EAGER_FLOATS_PER_ELEM} "
f"against fused {FUSED_FLOATS_PER_ELEM}")
print(" the two ratios are the whole prediction; the block also "
"contains a GEMM,")
print(" which moves the same bytes in both rows and drags the time "
"ratio down")
profile_block(eager, x, residual, "eager")
profile_block(fused, x, residual, "fused")
nvtx_pass(eager, fused, x, residual)
if not ok:
print("\nFAIL: the fused operator does not reproduce the eager chain")
return 1
print("\nPASS")
return 0
if __name__ == "__main__":
sys.exit(main())