Contents
Map

04 · Pretraining at Scale

Profile and Fuse

View as:

Code Lab 01 - Profile and Fuse

Profile one training step of a small GPT to see where the time goes, then take the memory-bound elementwise tail of a SwiGLU feed-forward block - rms_norm(silu(gate) * up + residual) * weight - and make it faster without changing the math: first with torch.compile, then with a hand-written Triton kernel. Every speedup is reported as a median with a bootstrap 95% confidence interval, because single timing runs lie.

← Back to Overview: Pretraining at Scale · Concepts: CUDA Concepts & GPU Profiling · Accelerators & Interconnects

Learning objectives 2 hours (plus GPU time for the Triton part)
By the end of this page you will be able to:
  • Profile a training step with torch.profiler and read the operator table and the Perfetto timeline
  • Estimate the minimum memory traffic of an elementwise chain and explain why fusing it helps
  • Fuse the chain with torch.compile and with a Triton kernel, checking numerical agreement with the eager version
  • Benchmark correctly - warm-up, synchronization, interleaved trials - and report speedups with bootstrap confidence intervals
Prerequisites

What's In This Lab

PropertyDetail
Part Atorch.profiler over 4 training steps of a 4-layer, d=256 GPT (RMSNorm, SDPA attention, SwiGLU FFN, AdamW); operator table plus trace.json for Perfetto
Part BThe fused tail on an 8192 x 4096 tensor: eager vs torch.compile vs Triton (CUDA only); max-difference check; 30 interleaved trials; bootstrap CI of the speedup
DevicesCUDA (BF16, all three variants) or CPU (FP32, eager and torch.compile)
VerifiedThe CPU path ran end to end on an Apple M4 laptop with PyTorch 2.14.1 in about 10 seconds (output below). The Triton/CUDA path was not run for this course - no GPU was available - so treat its numbers as yours to produce
Files01-Profile-and-Fuse/{profile_and_fuse.py, requirements.txt}
flowchart TD
    S["🏋️ Training step"] --> PR["🔬 torch.profiler<br/>operator table + trace"]
    PR --> H["🎯 Hot spot: elementwise tail<br/>silu · mul · add · rms_norm"]
    H --> E["🐢 Eager<br/>one kernel per op"]
    H --> C["⚙️ torch.compile<br/>Inductor fuses"]
    H --> T["🔧 Triton kernel<br/>one pass, registers only"]
    E & C & T --> B["📊 Interleaved trials<br/>median + bootstrap CI"]

    style H fill:#e8e0d4,stroke:#c8b89a
    style B fill:#dde4dc,stroke:#b0c4b0

Run It

cd src/content/04-Pretraining/CodeLabs/01-Profile-and-Fuse
python -m venv .venv && source .venv/bin/activate
pip install -r requirements.txt
python profile_and_fuse.py                     # cuda if available, else cpu
python profile_and_fuse.py --skip-profile --rows 16384 --cols 8192   # bigger tensor, Part B only

Verified CPU output (Apple M4, PyTorch 2.14.1, abbreviated):

=== Part A: top operators over 4 training steps (cpu) ===
                     Name    Self CPU %   Self CPU    # of Calls
                 aten::mm        52.54%   210.746ms          204
                aten::mul         7.28%    29.201ms          408
 aten::_scaled_dot_product_flash_atten...  7.06%   28.330ms   16
 aten::_log_softmax_backward_data  4.12%   16.515ms            4
       aten::_log_softmax         4.10%    16.448ms            4
      aten::silu_backward         2.36%     9.457ms           16
               aten::silu         2.29%     9.196ms           16
Self CPU time total: 401.088ms

                 eager: max |diff| vs eager = 0.00e+00
         torch.compile: max |diff| vs eager = 3.81e-06

=== Part B: fused tail, 8192 x 4096 float32 on cpu, 30 trials ===
minimum traffic if fused: 3 inputs + 1 output = 537 MB
                 eager: median  26.307 ms
         torch.compile: median  12.294 ms   speedup vs eager 2.14x  (95% CI 2.11-2.16)

On a CUDA GPU the script adds a triton (hand-written) row. Expect both fused variants to beat eager by a wider margin than on CPU, because GPU elementwise kernels are more purely bandwidth-bound and each eager launch also costs a few microseconds.


Walkthrough - What to Look At

  1. The operator table (Part A). Matrix multiplies (aten::mm) dominate, as they should in a healthy step - but mul, silu, softmax and their backward passes together take a large share for almost no FLOPs. Those are memory-bound, and they are the fusion opportunity. On CUDA, sort by CUDA time and look at the number of calls: hundreds of tiny kernels per step is the launch-overhead signature.
  2. Open trace.json in Perfetto. On a GPU, look for gaps between kernels (CPU-bound or synchronizing) and for long runs of short kernels (unfused elementwise work).
  3. The traffic estimate. The tail must read three tensors and write one: 537 MB for 8192 x 4096 FP32. Eager also writes and re-reads each intermediate (silu(gate), the product, h, h²...), several times that traffic. The fused versions approach the minimum.
  4. Correctness first. Every variant is compared with eager before timing. torch.compile differs by about 4e-6 - reordered floating-point arithmetic, not a bug. A fused kernel that is fast and wrong is worthless; always check.
  5. The Triton kernel (make_triton_tail) runs one program per row, because the RMS needs a reduction over the whole row. Everything between the loads and the store happens in registers - that is the fusion. It computes in FP32 internally even for BF16 inputs, which is why its output can match eager closely.
  6. The benchmark design. Warm-up runs exclude compilation; time_once uses CUDA events and synchronizes on GPU; variants are interleaved trial by trial so thermal or background drift affects all of them equally; the CI comes from resampling trials. A speedup whose CI includes 1.0 is not a speedup.

Check Yourself

Check yourself
0 / 4 answered
  1. In Part A, matrix multiplies take about half the CPU time, yet the lab targets the elementwise tail. Why?
  2. Why does the benchmark interleave the variants trial by trial rather than running all eager trials, then all compiled trials?
  3. torch.compile's output differs from eager by 3.8e-6. Is that a bug?
  4. Why does the Triton kernel use one program per row instead of splitting each row into many blocks?

Exercises

Exercise - Measure the fusion gain at different sizes

Run Part B at --rows 1024, 8192 and 32768 (cols 4096). How does the torch.compile speedup change with size on your hardware, and why?

Solution

At small sizes, fixed overheads - kernel launches on GPU, Python dispatch and thread start-up on CPU - are a large part of the time, and fusion's main benefit is fewer launches. At large sizes the work is dominated by memory traffic, and the speedup approaches the ratio of eager traffic to fused traffic. Plot speedup with its CI against size; if a small-size CI includes 1.0, report "no measurable gain" for that size rather than a number.

Exercise - Fuse the backward pass

The profile shows silu_backward and mul in the backward pass too. Wrap the tail in a small nn.Module, run forward and backward under torch.compile, and compare step time with eager. What does the compiler do with the backward?

Solution

torch.compile traces the backward through AOTAutograd and fuses it as well, so the backward of the tail (the gradients through RMSNorm, the residual add and silu * up) also becomes a few fused kernels. Measure forward plus backward with the same interleaved, bootstrapped method. For a hand-written Triton version you would need a second kernel for the backward and a torch.autograd.Function wrapping both - which is exactly the work the compiler saves you.

References

Last reviewed: 2026-10

⚡AI-assisted content - always verify, always explore multiple perspectives·