Optional torch.compile and TF32 for training - #407
Conversation
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
2d1be47 to
116a95f
Compare
…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
|
Hi @nihargupte-ph I can't find your glasflow branch? Did you push it? |
stephengreen
left a comment
There was a problem hiding this comment.
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: thetorch>=2.6bump is needed (test_epochenterseager_mode()unconditionally andtorch.compiler.set_stanceis 2.6+), butuv.lockwas not re-locked.compile_utils.py:48:torch_compile_cache_diris silently ignored in single-GPU mode (the cache setup is gated onrank is not None, andrun_trainingpasses no rank). Either honor it or say so in the docs.training_multi_gpu.md:220:Adagradacceptsfusedin 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_precisionis only recorded intrain_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, sofloat32_matmul_precision: highchanges 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.)
…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
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG
stephengreen
left a comment
There was a problem hiding this comment.
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.py37 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: Truefor 2 epochs, stage_1Falsefor 2 epochs) withtorch_compile: truekeepslayers_rbbit-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: highestis in the checkpoint metadata,dingo_lsreads 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
-
The glasflow fork vendors the wrong nflows commit.
submodules/nflowsinnihargupte-ph/glasflow@compile-friendly-rqspoints 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. (.gitmodulesnow names your fork, fine for now; revert touofgravity/nflowsonce merged.) -
The docs no longer say what compile needs, and the user gets no warning.
ad501acchad the "experimental until the static-shape spline is released" sentence and5ad09da7dropped it together with the fork note. With whatuv sync/pip install today (glasflow 0.4.1),torch_compile: truebreaks 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) Awarnings.warnincompile_networkwhen the installed spline lacks thecheck_domainkwarg (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 pinglasflow>=<release>. -
uofgravity/nflows#13 CI is red on its own new test.
test_matches_reference_implementationfails 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 assertstorch.equalbetween 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 oftorch.equalshould 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 unseededtest_forward_inverse_are_consistenttests also fail about one run in three with--rerunsdisabled).
Smaller
reset_graphs_if_requires_grad_changesalso 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_precisionoverwrites 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",inpyproject.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
left a comment
There was a problem hiding this comment.
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.
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>
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: truewraps the network withtorch.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: highenables TensorFloat-32 matmuls (10-bit-mantissa inputs, fp32 accumulation, everything else stays fp32).fused: trueunder the optimizer settings selects the fused Adam kernel; that needed no code, only docs.glasflow dependency.
torch.compileonly pays off if the rational-quadratic spline inglasflow.nflowsis 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: truestill runs but is slower than eager. A numerically identical static-shape rewrite is innihargupte-ph/nflows@compile-friendly-rqs, vendored bynihargupte-ph/glasflow@compile-friendly-rqs, and I will open PRs touofgravityfor both. Until that is released, usepyproject.tomlis unchanged apart fromtorch>=2.6.uv.lockis re-locked accordingly.Measurements
Production settings: the
npe_modelnetwork (hidden_dim1024, 30 flow steps), 10M SEOBNRv5PHM waveforms, per-GPU batch 4096, A100. "GPU step" is the single-GPU network time without dataloader; "realised" is a realdingo_trainepoch. All rows are fp32 apart from the last, i.e. no AMP.hidden_dim512, where the network is more launch-bound).automatic_mixed_precision: Truethe matmuls already run in fp16, sofloat32_matmul_precisionchanges 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.torch_compilepays off for trainings of tens of epochs or more.Review pointers
dingo/core/nn/compile_utils.pyis the whole compile mechanism.unwrap_network/get_ddp_moduleintorchutils.pymake checkpointing andno_sync()see through the compile wrapper.tests/core/test_compile.pyexercises 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