COURSE / SOURCE

profile_model.py

All lessons
Source filecode/day88-torch-profiler/profile_model.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 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())