Skip to main content

The Linter for PyTorch: Detects silent training bugs.

Project description

🔥 torch-audit

Runtime auditing for PyTorch training loops

PyPI License: MIT Python 3.10+ Code Style: Black CI

torch-audit is a “check engine light” for your training loop.

Unlike a static linter, torch-audit runs at runtime and inspects what actually happens during training:

  • real tensors and batches (device placement, suspicious ranges, layouts)
  • real optimizer configuration (weight decay pitfalls)
  • real gradients (NaNs/Infs, explosions, missing grads)
  • real model execution (unused “zombie” layers, stateful layer reuse)

The goal is to catch silent bugs that don’t crash your code but quietly ruin training or waste compute.


📦 Installation

pip install torch-audit

To run the optional integration demos you may also want:

pip install lightning transformers accelerate

🚀 Quick Start

Zero-touch mode: autopatch()

If you want the least code churn, use autopatch(). It monkey-patches model.forward and optimizer.step so a normal training loop automatically emits findings.

import torch
from torch_audit import autopatch
from torch_audit.reporters.console import ConsoleReporter

device = "cuda" if torch.cuda.is_available() else "cpu"
model = model.to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)

# Audits every 1000 optimizer steps (set to 1 to audit every step)
auditor = autopatch(
    model,
    optimizer=optimizer,
    every_n_steps=1000,
    reporters=[ConsoleReporter()],
    fail_level="ERROR",
    run_static=True,
    run_init=True,
)

for batch, targets in dataloader:
    batch = batch.to(device)
    targets = targets.to(device)

    optimizer.zero_grad(set_to_none=True)

    # ✅ No wrappers required
    outputs = model(batch)
    loss = criterion(outputs, targets)
    loss.backward()
    optimizer.step()

# Finalize results (and report again if you want)
result = auditor.finish(report=True)

# Restore original methods / detach hooks
# (recommended if you keep using the model after auditing)
auditor.unpatch()

Note: autopatch() modifies objects in-place. If you rely on compilation/tracing tools (e.g. torch.compile, TorchScript), prefer the explicit wrapper mode below.

Wrapper mode: audit_dynamic(...) + phase wrappers

If you want the most accurate phase reporting (forward/backward/optimizer) and the clearest control, wrap your loop with audit_dynamic(...) and call the wrappers.

import torch
from torch_audit import audit_dynamic
from torch_audit.reporters.console import ConsoleReporter

device = "cuda" if torch.cuda.is_available() else "cpu"
model = model.to(device)

optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)

# Audits every 1000 optimizer steps (set to 1 to audit every step)
with audit_dynamic(
    model,
    optimizer=optimizer,
    every_n_steps=1000,
    reporters=[ConsoleReporter()],
    fail_level="ERROR",
) as auditor:
    for batch, targets in dataloader:
        batch = batch.to(device)
        targets = targets.to(device)

        optimizer.zero_grad(set_to_none=True)

        # Use the wrappers so audits run in the right phase.
        outputs = auditor.forward(batch)
        loss = criterion(outputs, targets)
        auditor.backward(loss)
        auditor.optimizer_step()

# Finalize results (and report again if you want)
result = auditor.finish(report=False)

A lighter-weight option: audit_step()

If you already have a training-step function and want a minimal, opt-in integration, you can use the decorator. It runs a best-effort post-step audit (optimizer phase).

from torch_audit import Auditor, audit_step
from torch_audit.reporters.console import ConsoleReporter

auditor = Auditor(model, optimizer=optimizer, every_n_steps=1000, reporters=[ConsoleReporter()])

@audit_step(auditor)
def train_step(batch, targets):
    optimizer.zero_grad(set_to_none=True)
    out = model(batch)
    loss = criterion(out, targets)
    loss.backward()
    optimizer.step()
    return loss

with auditor:
    auditor.audit_static()
    auditor.audit_init()
    for batch, targets in dataloader:
        train_step(batch, targets)

If you want the most accurate runtime results (graph/activation checks, precise phase reporting), prefer the explicit wrappers: auditor.forward(), auditor.backward(), auditor.optimizer_step().


📂 Runnable demos

The examples/ folder contains runnable scripts designed to trigger findings.

  • python examples/demo_general.py — plain PyTorch loop, end-to-end runtime auditing
  • python examples/demo_cv.py — CV-ish model + common data/layout mistakes
  • python examples/demo_nlp.py — “NLP-ish” tensors (e.g., invalid token ids) + optimizer pitfalls
  • python examples/demo_lightning.py — Lightning integration pattern (demo includes a minimal callback)
  • python examples/demo_hf.py — Transformers pattern (no downloads; constructs a tiny model from config)
  • python examples/demo_accelerate.py — Accelerate pattern (audits around accelerator.backward(loss))

Note: the repository currently focuses on the core runtime engine. Some ecosystem integrations are shown as copy-paste patterns in the demos rather than shipped as a first-class API.


📚 Reference


🧰 One-shot audits (CLI / CI)

You can run an audit without a training loop (useful for CI smoke checks).

Tip: if the torch-audit command isn’t available in your environment, run the same command as: python -m torch_audit ...

# List all available rules
torch-audit --list-rules

# Explain a single rule (ID)
torch-audit --explain TA405

# Static checks (architecture / hardware hints)
torch-audit my_project.models:MyModel --phase static

# Init checks (optimizer config, weight decay pitfalls)
torch-audit my_project.models:MyModel --phase init

# JSON output (machine readable)
torch-audit my_project.models:MyModel --phase static -f json -o audit.json

# SARIF output (GitHub code scanning / security tab)
torch-audit my_project.models:MyModel --phase static -f sarif -o audit.sarif

Baselines and rule filtering

# Create / update a baseline file from the current findings
torch-audit my_project.models:MyModel --phase static --baseline baseline.json --update-baseline

# Only fail on new findings compared to the baseline
torch-audit my_project.models:MyModel --phase static --baseline baseline.json

# Run only specific rules
torch-audit my_project.models:MyModel --phase static --select TA200,TA202

# Ignore specific rules
torch-audit my_project.models:MyModel --phase static --ignore TA201

🛠️ What it checks today

This repo currently ships the following built-in validators:

Data integrity (runtime)

  • TA300 input device mismatch (e.g., CPU batch with GPU model)
  • TA301 suspicious float ranges (e.g., normalized data missing)
  • TA302 flat/empty tensors (near-zero variance)
  • TA303 suspicious layout heuristic (NHWC vs NCHW)
  • TA304 tiny batch sizes with BatchNorm
  • TA305 invalid integer inputs (e.g., negative token ids for embeddings)

Stability (runtime)

  • TA100 NaNs/Infs in parameters or gradients
  • TA102 gradient explosion (global grad norm)
  • TA103 “dead units” (exactly zero grads)
  • TA104 no gradients found
  • TA105 activation collapse / high sparsity (forward hooks)

Optimization config (static/init)

  • TA401 Adam + weight_decay (suggest AdamW)
  • TA402 weight decay applied to norm/bias params
  • TA403 weight decay applied to embeddings

Architecture + execution (static + runtime)

  • TA400 redundant bias before normalization
  • TA404 even convolution kernel sizes
  • TA405 dead convolution filters
  • TA500 unused “zombie” layers (runtime, forward hooks)
  • TA501 stateful layer reuse (e.g., BatchNorm called multiple times)

Hardware/performance hints (static/init)

  • TA200 tensor-core alignment hints
  • TA201 channels-last memory layout hints
  • TA202 model device placement / split-brain
  • TA203 AMP/precision suggestion

🧾 Reporters

You can output results to multiple formats:

from torch_audit.runtime import Auditor
from torch_audit.reporters.console import ConsoleReporter
from torch_audit.reporters.json import JSONReporter
from torch_audit.reporters.sarif import SARIFReporter

auditor = Auditor(
    model,
    optimizer=optimizer,
    reporters=[
        ConsoleReporter(),
        JSONReporter(dest="audit.json"),
        SARIFReporter(dest="audit.sarif"),
    ],
)

🤝 Contributing & feedback

If you find a silent bug torch-audit missed, or want a new runtime validator, please open an issue.

License

Distributed under the MIT License.

Project details


Download files

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

Source Distribution

torch_audit-0.3.0.tar.gz (36.6 kB view details)

Uploaded Source

Built Distribution

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

torch_audit-0.3.0-py3-none-any.whl (46.8 kB view details)

Uploaded Python 3

File details

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

File metadata

  • Download URL: torch_audit-0.3.0.tar.gz
  • Upload date:
  • Size: 36.6 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: poetry/2.1.1 CPython/3.11.4 Windows/10

File hashes

Hashes for torch_audit-0.3.0.tar.gz
Algorithm Hash digest
SHA256 8086ba5c76d10fd72532990aae18a81a1e5240c626821963f4e84e0afeed8873
MD5 71af5b970c05094c56a28bd111554fb3
BLAKE2b-256 cb273921c36c9ce90d596b9d66fd158c627de54a91166729723e910051acd70e

See more details on using hashes here.

File details

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

File metadata

  • Download URL: torch_audit-0.3.0-py3-none-any.whl
  • Upload date:
  • Size: 46.8 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: poetry/2.1.1 CPython/3.11.4 Windows/10

File hashes

Hashes for torch_audit-0.3.0-py3-none-any.whl
Algorithm Hash digest
SHA256 5386266299a700bfea383ccb1a3686bd6efd13ce3691f621d89a938041c435bc
MD5 7e2974eb83783b6eebf786d554d00a20
BLAKE2b-256 bb7f8b101807a3a0d702f7bc38b203fd6fb6e9617be1d53a08254d6ac36fbf41

See more details on using hashes here.

Supported by

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