TorchScan inspects a PyTorch model and returns a JSON-serializable report of its structure, parameters, inputs, module estimates, and operator FLOPs. Every metric says whether it is complete, partial, or unavailable, so an unsupported operation cannot masquerade as zero.
import torch.nn as nn
from torchscan import crawl_module, summary
model = nn.Conv2d(3, 8, 3)
# Print the human-readable table and receive the same structured report.
report = summary(model, (3, 32, 32))
# Or collect the report without printing the table.
report = crawl_module(model, (3, 32, 32), strict=True)summary keeps the familiar terminal UX while returning the structured report:
__________________________________________________________
Layer Type Output Shape Param # Trainable
==========================================================
conv2d Conv2d (1, 8, 30, 30) 224 True
==========================================================
Trainable params: 224
Non-trainable params: 0
Total params: 224
----------------------------------------------------------
Model size (params + buffers): 0.00 Mb
----------------------------------------------------------
Module-formula forward FLOPs: 388.80 kFLOPs
Multiply-Accumulations: 194.40 kMACs
Direct memory accesses: 201.82 kDMAs
Operator forward FLOPs: 388.80 kFLOPs
__________________________________________________________
input_shape excludes the batch dimension. For realistic calls—including masks, scalars, None, and nested
containers—pass complete args and kwargs instead:
import json
import torch
from torch import nn
from torchscan import crawl_module
class MaskedModel(nn.Module):
def forward(self, input_ids, *, attention_mask):
return input_ids * attention_mask
transformer_model = MaskedModel()
input_ids = torch.ones(1, 4)
attention_mask = torch.tensor([[True, True, False, False]])
report = crawl_module(
transformer_model,
args=(input_ids,),
kwargs={"attention_mask": attention_mask},
)
print(json.dumps(report["inputs"]["kwargs"]["attention_mask"], indent=2))Only metadata is retained:
{
"kind": "tensor",
"shape": [1, 4],
"dtype": "torch.bool",
"device": "cpu",
"requires_grad": false
}TorchScan temporarily evaluates the model with gradients disabled and restores every module's original training state. It records input metadata, never tensor values.
For shapes and parameter counts, skip compute analysis with one option:
report = summary(model, (3, 32, 32), mode="structure")Structure mode collects the same hierarchy, calls, input/output metadata, parameters, and buffers without FLOP
dispatch or module formulas. Unrequested compute totals have status="unavailable" and method="not_requested".
strict=True checks the requested metrics. Full analysis remains the default. Both modes release intermediate
activations as execution progresses.
Add your own module/model estimates with
custom_modules={YourModule: ModuleHandler(your_callback)} on crawl_module or summary. Callbacks receive the
complete actual call and supply FLOPs, MACs, DMAs, or receptive-field fields independently. Explicit subtree ownership
prevents inclusive parent estimates from double-counting children. Both APIs also accept custom_mapping for separate
operator FLOP overrides. Registrations belong to one analysis and require no TorchScan dependency on your model library.
See the copyable extension tutorial, including a complex-valued example and counting conventions.
Use zero-argument callables when the owner needs full control over execution:
import json
import torch
from torchscan import measure_flops
from torchscan.process import measure_peak_memory
inputs = torch.ones(8)
flops = measure_flops(lambda: torch.sin(inputs))
print(json.dumps(flops["total"], indent=2))
print("uncounted operator:", flops["diagnostics"][0]["operator"])
memory = measure_peak_memory(lambda: torch.cos(inputs), device=inputs.device)
print(memory["device"], memory["metric"])measure_flops uses PyTorch's operator dispatch. measure_peak_memory invokes the workload exactly once and reports
backend-specific PyTorch memory—not process RSS or total device memory.
Here, PyTorch has no built-in aten.sin formula, so TorchScan shows a lower bound instead of a false zero:
{
"status": "partial",
"value": null,
"known_value": 0,
"unit": "FLOPs",
"scope": "workload",
"method": "torch.utils.flop_counter.FlopCounterMode"
}
uncounted operator: aten.sin
cpu pytorch_tensor_bytes
Peak byte values are intentionally omitted because they depend on the workload, allocator, PyTorch version, and
hardware; the returned mapping includes baseline_bytes, peak_bytes, and delta_bytes.
import torch.nn as nn
from torchscan import compare_reports, crawl_module
before = crawl_module(nn.Conv2d(3, 8, 3), (3, 32, 32))
after = crawl_module(nn.Conv2d(3, 12, 3), (3, 32, 32))
diff = compare_reports(before, after)
parameters = diff["totals"]["parameters"]
print(parameters["status"], parameters["delta"])complete 112
compare_reports propagates incomplete metrics. It does not store baselines or decide whether a model fits a budget;
the model owner supplies those policies.
from pathlib import Path
import webbrowser
from torchscan import render_report
path = Path("torchscan-report.html").resolve()
path.write_text(render_report(report), encoding="utf-8")
webbrowser.open(path.as_uri())
# A static SVG, or an HTML comparison using compare_reports internally:
Path("torchscan-report.svg").write_text(render_report(report, format="svg"), encoding="utf-8")
Path("comparison.html").write_text(render_report(after, before=before), encoding="utf-8")HTML opens a module cost explorer: nested rectangles show the hierarchy and the concentration of recorded compute or first-attributed parameters. Select a module for tensor shapes, repeated-call evidence, methods, and diagnostics. An unscaled rail keeps unknown work and tiny/zero contributions visible; comparisons share one hierarchy and scale. SVG exports the same visual composition. Reports work offline with no server, CDN, or extra dependencies. See the guide for interpretation, comparison rules, and keyboard controls. Generate the local-model examples to explore incomplete work and a channel comparison.
complete: the requested scope was counted;valueis authoritative for the documented method.partial:known_valueis a lower bound and diagnostics identify missing work.unavailable: TorchScan cannot produce the metric for this execution.
Use strict=True when any incomplete analysis must stop automation. See the
report schema and
methodology before comparing results.
The stable release is v0.2.0. It requires Python ≥3.11,<4 and PyTorch ≥2.1,<3:
pip install torchscanThe unreleased version on main adds render_report, the custom_modules extension API, custom_mapping on
crawl_module and summary, and native Transformer MAC, DMA, and token-dependency estimates.
Install main to use these features:
pip install git+https://github.com/frgfm/torch-scan.gitSee the installation guide and v0.2 migration guide. For a local development checkout, follow Contributing.
- Agent quickstart
- Model and input support
- Custom module extensions
- v0.2 migration guide
- API reference
Agents can also load the repository skill at .agents/skills/torchscan/SKILL.md.
Citation metadata is available in CITATION.cff.
Contributions are welcome; see CONTRIBUTING.md. TorchScan is distributed under the
Apache License 2.0.
