Skip to main content

torch-preflight logo

torch-preflight

The linter that understands autograd — and the VRAM estimator that tells you what will fit.

CI PyPI Python versions License

Docs · Rules · VRAM estimation · CLI

A static analyzer for PyTorch training code. It catches VRAM leaks and silent convergence bugs at commit time, projects peak memory for a training run or a serving deployment before you rent the GPU, and can fail the run at step 0 instead of OOMing at step 400.

No GPU. No PyTorch. It never imports or executes your code.

$ torch-preflight check train.py

train.py
  7:19  error   TG001 (CRITICAL_OOM)
  `losses.append(...)` stores a tensor that is still attached to the autograd graph;
  every iteration's graph is retained in VRAM.
    7 │     losses.append(loss)
      │                   ^^^^
  help: Use `.item()` to keep just the scalar value, or `.detach()` to keep the tensor
        without its graph.
  fix:  add .detach() (run with --fix)

Found 1 error in 1 file(s).

One line, one wasted GPU hour. Caught in milliseconds, before it runs.

  • 🔍 13 rules for bugs ruff and flake8 cannot see — retained autograd graphs, missing zero_grad(), unscaled gradient accumulation, DDP without a DistributedSampler, a model stuck in eval() after validation, doubled softmax, unseeded runs, GPU sync in a hot loop
  • 🧮 VRAM estimation that covers the whole job — training, encoder-decoder, and autoregressive generation with the KV cache, plus the change that would make it fit
  • 🛡️ VRAMGuard — fails a run at step 0 rather than OOM at step 400, measuring your model's real activation memory without allocating a byte
  • 🛠️ Autofixes via concrete syntax tree rewrites, so formatting and comments survive untouched
  • 📊 Measured, not guessed — every constant calibrated against real hardware, 3.7% mean error versus measured peaks
  • 🤫 Quiet on real code23 findings across PyTorch's own 2,285 files, roughly one per hundred, every one triaged
  • No GPU and no PyTorch required — pure static analysis over LibCST; a CI job asserts torch is never imported
  • 🐍 Python 3.9–3.13, pyproject.toml config, pre-commit hook, GitHub Action, SARIF output

ruff and flake8 understand Python. They don't understand autograd graphs, gradient accumulation, or what num_workers=0 does to eight GPUs waiting on one CPU. torch-preflight is built for the bugs that only cost money once you're paying for a GPU.

Table of contents

Getting started

pip install torch-preflight
torch-preflight check ./src/                        # lint a tree
torch-preflight check ./src/ --fix                  # apply the safe fixes
torch-preflight estimate train.py --gpu a100-80gb   # will this training run fit?
torch-preflight estimate --model llama-3-8b --gpu a100-80gb \
    --generate --batch-size 16 --max-context 8192   # will this deployment fit?
torch-preflight explain TG003                       # why a rule exists, and what it costs

The base install has no heavy dependencies. torch-preflight[hub] adds Hugging Face architecture lookup; torch-preflight[vram] adds exact meta-device profiling.

The line that costs you a GPU hour

losses = []
for batch, targets in loader:
    optimizer.zero_grad()
    loss = criterion(model(batch), targets)
    loss.backward()
    optimizer.step()
    losses.append(loss)          # ← keeps every step's graph alive in VRAM

You have written this. Everyone has. loss still carries its computational graph, so appending it retains every intermediate activation from that step — and the next, and the next. Memory climbs linearly until CUDA gives up, hours in.

Why this is hard: losses.append(x) is only a bug when x carries a graph. torch-preflight runs a dataflow pass to find out, tracing values across assignments, arithmetic, tensor methods and function scopes, and refusing to propagate through .detach(), .item() or argmax. So losses.append(loss.item()) stays silent, and so does anything inside torch.no_grad(). A linter that pattern-matched on .append( would be unusable.

See all 13 rules →

Will this fit on the GPU I'm about to rent?

$ torch-preflight estimate finetune.py --gpu a100-80gb

Model      llama-2-7b  (arch-snapshot)   6.74 B params
Config     amp · AdamW · batch 4 · seq 2048
Read from  batch_size=finetune.py:9, optimizer=finetune.py:7, seq_len=finetune.py:13

  weights            25.10 GiB  ███
  gradients          25.10 GiB  ███
  optimizer state    50.21 GiB  ██████
  autocast cache     12.55 GiB  ██
  activations        67.42 GiB  ████████
  CUDA context         135 MiB  █
  fragmentation      18.94 GiB  ██
  ──────────────────────────────────────────────
  projected peak    199.45 GiB   (179.51 GiB – 219.40 GiB)

Target     NVIDIA A100 80GB (78.0 GiB usable)   →   256% of capacity   ✗ OOM

What would make it fit:
  ✗  − 35.36 GiB  →  164.09 GiB   flash attention / SDPA
       mathematically equivalent, removes the O(seq²) attention term
  ✗  − 66.30 GiB  →  133.15 GiB   gradient checkpointing
       same result, roughly 30% slower
  ✗  − 41.61 GiB  →  157.84 GiB   8-bit AdamW (bitsandbytes)
       quantised optimizer state, minimal quality impact
  ✗  −112.56 GiB  →   86.89 GiB   flash attention + checkpointing + 8-bit AdamW
                                  + halve micro-batch
       even stacked together these do not fit; this needs a larger GPU, more
       devices, or a parameter-efficient method such as LoRA

A single 80GB A100 is the wrong tool for a full 7B fine-tune at sequence 2048. Better to learn that now than after the instance is running.

Every field is read out of your script — model, batch size, sequence length, precision, optimizer, sharding — and the report says which line each came from, so you can check it. Nothing is imported or executed. 41 architectures ship built in, 23 GPUs and 34 cloud instances are known by name (--gpu p4de.24xlarge works), and anything else is measured exactly on PyTorch's meta device without allocating a byte.

Every other estimator stops at the number. The list of what to change is the part you actually wanted.

Serving, not just training

Generation is a different memory shape, and --generate models it:

$ torch-preflight estimate --model llama-3-8b --gpu a100-80gb \
      --generate --batch-size 16 --max-context 8192 --dtype pure-bf16

Config     pure-bf16 · batch 16 · context 8192 · generation (KV cache)

  weights            14.96 GiB  ███████████
  activations           20 MiB  █
  KV cache           16.00 GiB  ████████████
  ...
  projected peak     32.68 GiB   (29.41 GiB – 35.95 GiB)

Target     NVIDIA A100 80GB (78.0 GiB usable)   →   42% of capacity   ✓ FITS

The KV cache is usually the term that decides the answer, and it is where grouped-query attention pays off: llama-3-8b's 8 KV heads against 32 query heads make its cache a quarter of the multi-head equivalent. Llama-2-7b, which has none, needs 64 GiB of cache against 12.6 GiB of weights at batch 32 and 4096 tokens.

Decoding also collapses the activation term — one token attends against the cache, so the O(context²) attention matrix never materialises.

Fail at step 0, not step 400

from torch_preflight import VRAMGuard

with VRAMGuard(model, optimizer=optimizer, batch_size=32, image_size=224):
    train()

Parameters, gradients and optimizer state come from the live model and are exact. Activations are measured from your module, by running one forward pass against meta-device parameters — that allocates nothing, never touches your model, and costs milliseconds. It raises only when the run cannot fit even at the optimistic end of the interval; aborting a job on a guess would be worse than the OOM.

It reads your config, not just your code

DeepSpeed ZeRO stage and CPU offload are parsed out of the JSON or dict your script points at — reading JSON is not executing code — so a ZeRO-3 run with offload_optimizer is charged for what actually sits on the device. T5 and Whisper are modelled as encoder-decoders, with the cross-attention term a decoder-only formula cannot express.

See VRAM estimation →

Why you can trust the numbers

Memory estimators are easy to write and easy to be quietly wrong about. So:

Constants are measured Activation coefficients from saved_tensors_hooks on the meta device; allocator behaviour and CUDA context from a real GPU. Measurement showed the published Megatron constants are a midpoint of two regimes — models with dropout retain 3× the attention tensors — so Llama-class models are charged the cheaper rate they actually pay.
Projections are checked 3.7% mean absolute error against measured peaks for GPT-2, BERT, DistilBERT and ResNet-50 on a T4. Harness and fixtures in tests/calibration/, so you can re-run them.
Gaps are stated, not papered over offload_param streams parameters in, so the resident set is smaller than the weights term — we have not measured it, so the report says the real peak is lower rather than inventing a fraction. Grouped-query attention does not reduce training activations (measured: repeat_kv materialises full-size K/V), so it is not modelled as if it did.
It refuses to guess An unrecognised model reports UNKNOWN and widens the interval rather than inventing a parameter count. Verdicts are bands with an error range, never a fabricated "95% risk" score.
It stays quiet 23 findings across PyTorch's own 2,285 files — about one per hundred — every one triaged as deliberate. Those passes found real bugs in the rules, each now regression-tested.

416 tests. A typical project lints in well under a second; PyTorch's entire 2,285-file source tree takes about four minutes with all 13 rules.

Integrations

# .pre-commit-config.yaml
repos:
  - repo: https://github.com/highwaterlabs/torch-preflight
    rev: v0.2.0
    hooks:
      - id: torch-preflight
# .github/workflows/lint.yml
- uses: highwaterlabs/torch-preflight@v0
  with:
    paths: src/
    format: github      # inline PR annotations

SARIF output feeds GitHub code scanning; JSON feeds everything else. Set target_gpu in pyproject.toml and CI fails on a projected OOM before the job is ever submitted.

See CI integration →

Documentation

Rules All 13 rules, and the false positives deliberately suppressed
VRAM estimation Custom architectures, CI gating, VRAMGuard, accuracy
CLI reference Commands, flags, exit codes, autofixes
Configuration pyproject.toml and inline suppression
CI integration GitHub Action, pre-commit, SARIF
Architecture How the analysis pipeline works
Development Tests, adding a rule, roadmap

Design notes live in design/, including the RFC behind the estimator and the spike the cost model rests on.

What stays free

MIT licensed. These are commitments, not just current state:

  • Every rule that has ever shipped free stays free.
  • The estimator, the remediation solver and VRAMGuard stay complete — not a demo tier.
  • The rule API stays open, so anyone can write and ship their own rules.
  • The calibration method and data stay public and reproducible. Numbers are only worth trusting if you can check them.

A hosted service may come later for things that genuinely need a server or a team. Nothing above is part of that.

Contributing

Issues and pull requests are welcome. Adding a rule is one file plus a @register decorator — see development for the walkthrough and the test conventions.

License

MIT — see LICENSE.

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

torch_preflight-0.3.0.tar.gz (206.1 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

torch_preflight-0.3.0-py3-none-any.whl (130.3 kB view details)

Uploaded Python 3

File details

Details for the file torch_preflight-0.3.0.tar.gz.

File metadata

  • Download URL: torch_preflight-0.3.0.tar.gz
  • Upload date:
  • Size: 206.1 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for torch_preflight-0.3.0.tar.gz
Algorithm Hash digest
SHA256 a8f06c85595babf4f7186c3eb881a9c0b48fcb6807ba3aa6d0dd132680973546
MD5 96077b5d2e58a6eccaa0924abb390042
BLAKE2b-256 0e9eeaff0a911a1d42d2d563ce9d6b5d593ecb22e7f87e9892e2274f9d71f584

See more details on using hashes here.

Provenance

The following attestation bundles were made for torch_preflight-0.3.0.tar.gz:

Publisher: release.yml on highwaterlabs/torch-preflight

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file torch_preflight-0.3.0-py3-none-any.whl.

File metadata

  • Download URL: torch_preflight-0.3.0-py3-none-any.whl
  • Upload date:
  • Size: 130.3 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for torch_preflight-0.3.0-py3-none-any.whl
Algorithm Hash digest
SHA256 9fb7f11a828d50b1a7a50189a3a488fc7438a6fd7408943a81c01b9bc1db93ec
MD5 d021495f46c013987344f474cbfa0c3d
BLAKE2b-256 56213991c919bd41fca3d9a1556e42d403f579bea9b9ae51ecd9cbf7fbcd39dc

See more details on using hashes here.

Provenance

The following attestation bundles were made for torch_preflight-0.3.0-py3-none-any.whl:

Publisher: release.yml on highwaterlabs/torch-preflight

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page