Skip to content

Optional torch.compile and TF32 for training - #407

Merged
stephengreen merged 8 commits into
mainfrom
multi-gpu-optimizations
Sep 11, 2026
Merged

stephengreen merged 8 commits into
mainfrom
multi-gpu-optimizations

Conversation

@nihargupte-ph

@nihargupte-ph nihargupte-ph commented Sep 9, 2026 •

Copy link
Copy Markdown
Collaborator

Two opt-in speed settings for training, both off by default, on top of the DDP work in #358.

Why

Profiling the DDP scaling showed that the NPE network is not compute-bound but launch-bound: the spline flow issues tens of thousands of tiny CUDA kernels per step, and the CPU cannot dispatch them faster than the GPU finishes them. GPU utilization plateaus around 60% even at batch sizes that fill the memory. Separately, PyTorch runs float32 matmuls at full precision by default, so the tensor cores of A100-class GPUs sit idle.

What is added

local.torch_compile: true wraps the network with torch.compile, which fuses the small kernels. It works in the single-GPU and the DDP path; under DDP each rank compiles into its own node-local cache (local.torch_compile_cache_dir, optional). The network is compiled once per run: the last partial batch of each epoch is dropped and the test epoch runs eagerly, since either would otherwise trigger a full recompilation. Checkpoints of compiled networks save with the usual keys and load anywhere.

local.float32_matmul_precision: high enables TensorFloat-32 matmuls (10-bit-mantissa inputs, fp32 accumulation, everything else stays fp32). fused: true under the optimizer settings selects the fused Adam kernel; that needed no code, only docs.

glasflow dependency. torch.compile only pays off if the rational-quadratic spline in glasflow.nflows is written with static shapes. The released version selects the spline tails with mask indexing, which breaks the compiled graph in every transform; with it, torch_compile: true still runs but is slower than eager. A numerically identical static-shape rewrite is in nihargupte-ph/nflows@compile-friendly-rqs, vendored by nihargupte-ph/glasflow@compile-friendly-rqs, and I will open PRs to uofgravity for both. Until that is released, use

pip install git+https://github.com/nihargupte-ph/glasflow@compile-friendly-rqs

pyproject.toml is unchanged apart from torch>=2.6. uv.lock is re-locked accordingly.

Measurements

Production settings: the npe_model network (hidden_dim 1024, 30 flow steps), 10M SEOBNRv5PHM waveforms, per-GPU batch 4096, A100. "GPU step" is the single-GPU network time without dataloader; "realised" is a real dingo_train epoch. All rows are fp32 apart from the last, i.e. no AMP.

configuration GPU step realised step, 1 GPU realised step, 4 GPUs
fp32, no compile (today) 0.85 s 0.85 s 0.88 s
compile only 0.70 s 0.70 s 0.77 s
compile + TF32 + fused Adam 0.28 s 0.40 s 0.62 s
compile + fp16 AMP + fused Adam 0.21 s (bf16) – 0.63 s
  • Compile alone is 1.2x at this width (1.4x at hidden_dim 512, where the network is more launch-bound).
  • All three together make the GPU step 3x faster (0.85 to 0.28 s). The realised gain is smaller, 2.1x on 1 GPU and 1.4x on 4 GPUs, because the CPU data pipeline then becomes the limit; that is a separate problem, not addressed here.
  • TF32 and AMP are alternatives, not a stack. Under automatic_mixed_precision: True the matmuls already run in fp16, so float32_matmul_precision changes nothing; the 3x above is the no-AMP stack. TF32 alone gives the same 1.9x as AMP alone on this network (0.90 to 0.47 s per step), and is the option for trainings that stay in fp32. With AMP, the same stack (compile + AMP + fused Adam) is as fast or slightly faster.
  • Accuracy: 3-epoch trainings with fp32, TF32 and fp16 AMP agree in train and test loss to within run-to-run noise (differences of 0.01 or less). TF32 stays opt-in; a longer comparison is advisable before making it a default.
  • Compilation costs 5 to 12 minutes once per run and at every stage boundary that changes the trainable parameters, so torch_compile pays off for trainings of tens of epochs or more.

Review pointers

  • dingo/core/nn/compile_utils.py is the whole compile mechanism.
  • unwrap_network / get_ddp_module in torchutils.py make checkpointing and no_sync() see through the compile wrapper.
  • tests/core/test_compile.py exercises the compile path; the no-graph-break test only passes with the glasflow fork installed.

🤖 Generated with Claude Code

https://claude.ai/code/session_0119JTh4nB6zdJNXE7tPBaLN

nihargupte-ph and others added 4 commits September 9, 2026 22:05
The neural spline flow launches tens of thousands of tiny CUDA kernels per
step, so on an A100 the step is bound by kernel-launch overhead rather than
by arithmetic (GPU utilization plateaus at ~60%, even at batch sizes that
fill the memory). torch.compile fuses those kernels and removes most of the
overhead, but only if the rational-quadratic spline in glasflow.nflows is
written with static shapes; with the released implementation (boolean-mask
indexing and a host-syncing branch) it graph-breaks in every coupling
transform and is a net slowdown.

Add an opt-in local.torch_compile setting (default false) that wraps the
network with torch.compile, in both the single-GPU and the DDP path (after
the DDP wrap, so the gradient all-reduce keeps overlapping with backward).
Under DDP each rank gets its own node-local Inductor/Triton cache
(local.torch_compile_cache_dir), since concurrent ranks race on the shared
cache and Triton artifacts on a network filesystem can be unloadable.

The static-shape spline lives in the compile-friendly-rqs branches of
nihargupte-ph/nflows and nihargupte-ph/glasflow until it is merged upstream;
compile_network raises with the install command if the installed glasflow
lacks it. Checkpoints and the DDP no_sync() path now look through the
torch.compile wrapper (unwrap_network / get_ddp_module), so compiled
networks save without wrapper prefixes and load anywhere.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01JXDcmhZcz96xz9Zm475iAa
With torch_compile enabled a production-size run compiled the network four
times in its first epoch (5-14 minutes each): the training graph, a second
graph for the smaller last batch of the epoch (the compiled graph is
specialized to the batch shape), and both again for the test epoch in eval
mode. That made the first epoch slower than eager and pushed break-even to
6-9 epochs.

Drop the last partial batch of the train and test loaders when torch_compile
is set (drop_last on build_train_and_test_loaders), and run the test epoch
under torch.compiler.set_stance("force_eager") so the compiled network is
evaluated eagerly: a test epoch is ~5% of the training epoch and never
amortizes its own compilation.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VXADmfii2HyfuFTsGcWb1V
PyTorch keeps float32 matmuls at full precision by default, so the tensor
cores of A100-class GPUs go unused for the flow's linear layers. Expose
torch.set_float32_matmul_precision through a local setting (highest | high |
medium; default unchanged) applied in the single-GPU and every DDP process.
For the npe_model network at per-GPU batch 4096, 'high' roughly halves the
step time in a dataloader-free benchmark; accuracy is to be validated per
training before adopting it.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01JbypaGF6ZbFyMNJk52hLPF
… the fact that we upgraded torch dependancy to be >2.6
@nihargupte-ph
nihargupte-ph force-pushed the multi-gpu-optimizations branch from 2d1be47 to 116a95f Compare September 9, 2026 20:05
…not installed

Released glasflow selects the RQ-spline tails with mask indexing, so the
fullgraph trace fails on CI. The test now checks for the check_domain
kwarg that marks the compile-friendly-rqs fork and skips otherwise.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG
@stephengreen

Copy link
Copy Markdown
Member

Hi @nihargupte-ph I can't find your glasflow branch? Did you push it?

@stephengreen stephengreen left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks, the mechanism is clean: the wrapper unwrapping, the checkpoint keys, the eager test epoch and the compile-outside-DDP order all check out (details at the end). Three things need fixing before this can merge, plus a few smaller ones.

1. Unfreezing the RB layer at a stage boundary does not recompile; the RB layer silently never trains (blocking)

train_pipeline.py:305-311 flips requires_grad on the already-compiled pm.network (compiled once at :549 / :674, before train_stages). Dynamo does not guard on a parameter's requires_grad, so the graph traced in stage 0 with the RB layer frozen is reused in stage 1, and its backward produces no gradient for those parameters. The rebuilt optimizer skips them. No error, and the loss curve looks normal. This is exactly the examples/npe_model layout (stage_0: freeze_rb_layer: True, stage_1: False) with torch_compile: true. A resume from a checkpoint exactly at the boundary escapes it, because the fresh network is traced with the current flags.

Reproduced on CPU (torch 2.9.1, inductor and aot_eager), on the real FlowWrapper and on a bare nn.Sequential(Linear, Linear): after unfreezing, layers_rb.*.grad is None and zero new graphs are traced. Re-wrapping with a fresh torch.compile does not help (the cache is keyed on the forward code object and its guards still pass). torch.compiler.reset() before the first stage-1 forward restores the gradient.

Suggested fix: keep compile_network where it is and, in initialize_stage, call torch.compiler.reset() when torch_compile is on and the freeze flag changes between stages, before the first forward of the new stage. That is the recompile the docs already promise (training_multi_gpu.md:171). Please add a test that compiles a small network, flips the flag, and asserts a gradient appears.

2. The glasflow fork is not on GitHub, and CI is red (blocking)

nihargupte-ph/glasflow and nihargupte-ph/nflows have no compile-friendly-rqs ref (git ls-remote), both are identical to upstream, and there is no PR on uofgravity/glasflow or uofgravity/nflows. So the pip install git+... line in the description cannot work and the benchmark is not reproducible. With released glasflow 0.4.1, tests/core/test_compile.py:125 (test_flow_compiles_without_graph_breaks) fails rather than skips (Unsupported: Dynamic shape operator at rational_quadratic.py:38), on all four CI jobs.

Could you push the branch and open the upstream PRs, and make that test skip (or xfail) when the fork is not installed? Until glasflow releases the rewrite, I would mark torch_compile as experimental in the docs. One thing to watch in the static-shape spline: if the tails are selected with torch.where, an inf or NaN in the unselected branch still poisons the gradient, so the clamp has to happen before the spline arithmetic.

3. drop_last is also applied to the test loader

torchutils.py:445,467. The test epoch runs eagerly, so there is nothing to protect there. When a rank's test split is smaller than the batch size the test loader is empty and test_epoch returns NaN (0/0 in AvgTracker), which then feeds the scheduler, the history file and early stopping; under DDP the per-rank split is len/world_size, so this is easy to hit. Otherwise it just biases the test loss by up to batch_size - 1 samples. Train loader only, please.

Smaller

  • pyproject.toml: the torch>=2.6 bump is needed (test_epoch enters eager_mode() unconditionally and torch.compiler.set_stance is 2.6+), but uv.lock was not re-locked.
  • compile_utils.py:48: torch_compile_cache_dir is silently ignored in single-GPU mode (the cache setup is gated on rank is not None, and run_training passes no rank). Either honor it or say so in the docs.
  • training_multi_gpu.md:220: Adagrad accepts fused in torch 2.9, and fused Adam runs on CPU tensors too. The compile note should also say which loader drops its last batch.
  • float32_matmul_precision is only recorded in train_dir/local_settings.yaml (train_pipeline.py:898), not in the checkpoint metadata. Worth storing with the model, since it is a training-precision choice.
  • TF32 vs AMP: under automatic_mixed_precision, autocast already runs the matmuls in fp16, so float32_matmul_precision: high changes essentially nothing in that case; it is the speedup for runs without AMP. A sentence in the guide would save people from setting both and expecting them to stack.

Checked and fine

unwrap_network / get_ddp_module in both nesting orders; checkpoint round trip on CPU (plain keys, loads into an uncompiled model with identical log_prob, compiled resume from an uncompiled checkpoint, optimizer sees the live parameters); eager_mode (the eager test epoch traces nothing, while eval mode or an odd batch would each recompile, so both mitigations are real); DDP (gloo) + compile forward/backward with no_sync; DistributedSampler padding means every rank drops the same batch; TF32 is set in every DDP worker; fused passes through and sets _step_supports_amp_scaling, so scaler.step is fine.

(Environment: torch 2.9.1, CPU, released glasflow 0.4.1.)

nihargupte-ph and others added 3 commits September 11, 2026 11:19
…in loader only

- Dynamo does not guard on requires_grad of parameters, so a network
  compiled with the RB layer frozen kept reusing that graph after the
  layer was unfrozen at a stage boundary and never trained it. Reset the
  compiled graphs (torch.compiler.reset) in initialize_stage when the
  freeze flag changes the trainable set; tests for the failure mode and
  the fix, verified end to end on a two-stage tutorial training.
- drop_last now only applies to the training loader. The test epoch runs
  eagerly, and dropping the last test batch could empty a small per-rank
  split and feed NaN to the scheduler and early stopping.
- torch_compile_cache_dir is honored in single-GPU mode as well.
- Record float32_matmul_precision in the checkpoint metadata.
- Docs: torch_compile marked experimental until glasflow releases the
  static-shape spline; which loader drops its last batch; fused optimizer
  support (adagrad included, CPU supported); TF32 and AMP are alternatives.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG

@stephengreen stephengreen left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks @nihargupte-ph, this addresses everything from the first round, and the fork checks out. Three things to fix before merging, all small.

What I verified (macOS arm64 CPU, torch 2.9.1)

  • Released glasflow 0.4.1: tests/core/test_compile.py + test_multi_gpu.py 37 passed, 1 skipped (the fork-only test). CI green on all four Python jobs.
  • With the fork (pip install git+...glasflow@compile-friendly-rqs): 13 passed, 0 skipped; torch._dynamo.explain(flow.log_prob) gives 1 graph, 0 breaks (released: 17 graphs, 16 breaks).
  • Numerics of the static-shape spline vs released glasflow, both fork versions (63658c2 and 20a62661): log_prob, all parameter gradients, samples (inverse path), and the spline called directly (float32/float64, forward/inverse, gradients wrt inputs and all three parameter sets, ~70% of inputs in the tails) are bit-identical (max |Δ| = 0). Same for tensor shapes up to (4096, 15).
  • The RB-layer fix, end to end: a two-stage CPU training (the small NSF from the hackathon-1 integration fixture, 4000 IMRPhenomXPHM waveforms, GWOSC O1 ASDs; stage_0 freeze_rb_layer: True for 2 epochs, stage_1 False for 2 epochs) with torch_compile: true keeps layers_rb bit-identical across stage 0 and changes it in every stage-1 epoch (max |Δ| 3.6e-2 then 1.8e-2), the same as the eager control run. Test loss is finite every epoch (the 200-sample test split ran its final batch of 8), float32_matmul_precision: highest is in the checkpoint metadata, dingo_ls reads it, and the checkpoint loads with plain state-dict keys and samples. Dynamo recompiles: one group of 7 one-time retraces right after the first compile (requires_grad and masked-size guards of the released spline) and the same group once after the stage boundary, none per epoch.

To fix

  1. The glasflow fork vendors the wrong nflows commit. submodules/nflows in nihargupte-ph/glasflow@compile-friendly-rqs points at 63658c2 (clamp applied to all inputs), while the nflows branch and uofgravity/nflows#13 are at 20a62661 (torch.where(inside, inputs, clamp), your fix for the clamp subgradient at the bound). I checked the pip-installed file: byte-identical to 63658c2. Nothing goes wrong on torch 2.9.1 (bit-identical either way, see above), but the install line in the description doesn't deliver what is under review upstream. Please bump the pointer. (.gitmodules now names your fork, fine for now; revert to uofgravity/nflows once merged.)

  2. The docs no longer say what compile needs, and the user gets no warning. ad501acc had the "experimental until the static-shape spline is released" sentence and 5ad09da7 dropped it together with the fork note. With what uv sync/pip install today (glasflow 0.4.1), torch_compile: true breaks the graph in every coupling transform and by your own measurement is slower than eager, after a 5–12 minute compile. Two small things: (i) restore one sentence in the docs (no fork URL needed): experimental; requires the static-shape rational-quadratic spline (uofgravity/nflows#13, not yet released); with glasflow ≤ 0.4.1 it runs but is slower than eager. (ii) A warnings.warn in compile_network when the installed spline lacks the check_domain kwarg (the same probe as the test), naming the remedy, so the user is told at startup rather than after an hour. The option stays off by default; both go away once we can pin glasflow>=<release>.

  3. uofgravity/nflows#13 CI is red on its own new test. test_matches_reference_implementation fails on Ubuntu/3.11/torch 2.14.0 (float64, mixed, forward) and Windows/3.8/torch 2.4.1 ("logabsdet differ"), while passing on Ubuntu/3.10 with the same torch 2.14.0. The test asserts torch.equal between the spline evaluated on the full tensor and on the masked gather, and bit-identity across tensor sizes isn't something PyTorch's kernels guarantee (vectorized loops handle the remainder differently by ISA/build), which fits "same torch, different runner". I could not reproduce it here (all differences exactly zero), so it is likely last-bit. A tolerance (self.eps ~1e-6 for float32, ~1e-12 for float64) instead of torch.equal should make it robust; exact equality can stay for the tail rows if you want to keep that guarantee visible. The Ubuntu/3.10 failure is the pre-existing unseeded cubic-coupling test (your diff doesn't touch it; locally the unseeded test_forward_inverse_are_consistent tests also fail about one run in three with --reruns disabled).

Smaller

  • reset_graphs_if_requires_grad_changes also fires at the very start of stage_0 (fresh network, all parameters trainable, stage_0 freezes), so the log says "compiled graphs (if any) discarded" before anything has compiled. Harmless (reset on an empty cache); maybe reword to "trainable parameters change; (re)compiling on the next step".
  • _record_float32_matmul_precision overwrites on resume, so a run resumed under a different precision keeps only the latest value. Fine for now; worth knowing.
  • Trailing whitespace after "torch>=2.6", in pyproject.toml.

Checked and fine

torch.compiler.reset() placement (before set_requires_grad_flag, after the optimizer is built; the optimizer holds all parameters anyway, so the unfrozen ones pick up their gradients); the resume path (initialize_stage(resume=True) runs after compile_network but before any trace, so the reset is a no-op and the first trace sees the right flags); DDP (freezing raises; unfreezing after a single-GPU frozen stage goes through a fresh load with all flags True, so nothing to reset; every rank calls the process-global reset); is_compiled in both nesting orders; test loaders keep their last batch in both branches; single-GPU cache dir; uv.lock (torch specifier, htcondor bump, linux markers on the nvidia libs; torch stays 2.9.1); no consumer pins the checkpoint metadata key set (grepped here and on hackathon-1).

@stephengreen stephengreen left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Approving this and will merge. Just be aware:

The fork branch compile-friendly-rqs on your glasflow is unchanged since 2026-09-02 and its nflows submodule still points at the clamp-all commit, not the torch.where fix that is under review upstream. So you need to make sure you change it if you plan to use that branch.

@stephengreen
stephengreen merged commit 5b5c5f0 into main Sep 11, 2026
5 checks passed
stephengreen added a commit that referenced this pull request Sep 11, 2026
Brings in the torch.compile and TensorFloat-32 opt-ins (#407) and the
SVDBasis loading fix after the dtype_map change (#411).

Conflicts resolved: the import block of base_model.py (both sides added an
import next to the backward_compatibility line; kept both) and the torch
specifier line in uv.lock (kept >=2.6, matching the merged pyproject).

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants