Metadata-Version: 2.5
Name: torchtyc
Version: 0.3.1
Summary: Shape checking for PyTorch, from your jaxtyping annotations
Project-URL: Homepage, https://github.com/bhaswata08/torchtyc
Project-URL: Repository, https://github.com/bhaswata08/torchtyc
Project-URL: Issues, https://github.com/bhaswata08/torchtyc/issues
License: Apache-2.0
License-File: LICENSE
Requires-Python: >=3.11
Requires-Dist: jaxtyping>=0.3
Provides-Extra: all
Requires-Dist: einops>=0.8; extra == 'all'
Requires-Dist: pygls>=1.3; extra == 'all'
Requires-Dist: watchfiles>=0.24; extra == 'all'
Provides-Extra: einops
Requires-Dist: einops>=0.8; extra == 'einops'
Provides-Extra: lsp
Requires-Dist: pygls>=1.3; extra == 'lsp'
Provides-Extra: watch
Requires-Dist: watchfiles>=0.24; extra == 'watch'
Description-Content-Type: text/markdown

# torchtyc

Shape checking for PyTorch, from your jaxtyping annotations.

torchtyc reads [jaxtyping](https://docs.kidger.site/jaxtyping/) annotations and
verifies the shapes your code actually produces, before you run it on a GPU. It
is the PyTorch counterpart to [jaxtyc](https://github.com/BeeGass/jaxtyc), which
does the same thing for JAX with `jax.eval_shape`.

```python
# model.py
import torch
from einops import einsum
from jaxtyping import Float
from torch import Tensor, nn


class Linear(nn.Module):
    def __init__(self, in_features: int, out_features: int) -> None:
        super().__init__()
        self.W = nn.Parameter(torch.empty((out_features, in_features)))

    def forward(self, x: Float[Tensor, "... in_features"]) -> Float[Tensor, "... in_features"]:
        return einsum(x, self.W, "... in_features, out_features in_features -> ... out_features")
```

```
$ torchtyc check model.py
model.py:12:63: error[shape-mismatch]
  in the return of `Linear.forward`: annotated `in_features`, but the traced dimension is `out_features`
    Expected: (..., in_features)
    Got:      (..., out_features)
    12 | def forward(self, x: Float[Tensor, "... in_features"]) -> Float[Tensor, "... in_features"]:
  try:  Float[Tensor, "... out_features"]
  hint: this dimension is `out_features`, so the annotation likely names the wrong axis

Found 1 error(s) in 1 function(s) across 1 file(s)
```

Pyright and mypy cannot do this. To them `Float[Tensor, "... in_features"]` is
just `Tensor`, and the dim string is an opaque literal.

## How it works

torchtyc constructs each annotated function's arguments on
`torch.device("meta")`, calls the function, and compares the shape that comes
back against the annotation. Meta tensors carry shape, dtype, and stride but own
no storage, so a whole model runs for the cost of a dictionary lookup per
operator, with no allocation and no arithmetic.

Every dimension name is bound to a distinct prime number starting at 101. This
buys two things:

**Mistakes cannot hide.** If `d_in` and `d_out` were both bound to 64, a
transposed weight matrix would sail through. Distinct primes make the two
impossible to confuse.

**Products stay readable.** A traced dimension that is the product of two
bound primes factors back into the names that produced it, which tells you the
function flattened two axes together. The primes themselves never reach the
message:

```python
# flat.py
def flatten(x: Float[Tensor, "batch seq d_model"]) -> Float[Tensor, "batch seq d_model"]:
    return x.reshape(x.shape[0], -1)
```

```
flat.py:6:55: error[rank-mismatch]
  in the return of `flatten`: expected 3 dimensions, traced 2
    Expected: (batch, seq, d_model)
    Got:      (batch, d_model*seq)
    6 | def flatten(x: Float[Tensor, "batch seq d_model"]) -> Float[Tensor, "batch seq d_model"]:
  hint: d_model*seq looks like two annotated axes flattened into one
```

An axis you did not name renders as `...`, and a single `_` axis renders as `_`.
A width your own `__init__` worked out renders as what it follows, such as
`<from d_model>`, because the number it lands on comes from the width torchtyc
traced with and is not a width your model has. A plain number appears only for a
size your code writes down, which is one you can go and find.

Because it runs your function rather than reasoning about it symbolically,
torchtyc also catches anything that raises on the way: a bad `einsum`, a
`matmul` between incompatible operands, an `nn.Module` that cannot be built.

```
layers.py:129:45: error[trace-error]
  RuntimeError: einsum(): subscript a has size <from d_model> for operand 1 which does not broadcast with previously seen size d_model
    129 | w1_out: Float[Tensor, "... d_ff"] = self.w1(x)
  hint: raised further down, in `Linear.forward` at line 40, where x is (..., d_model), self.weight is (d_model, <from d_model>)
  note: <from d_model> is a width your __init__ computed from d_model, so it is not d_model
```

The diagnostic anchors to the line you wrote, not to a line in torch and not to
the shared layer further down that every caller runs. The hint carries the
shapes of the tensors on the line that raised, so both widths that disagree are
on show: here `self.w1` was built with its two widths the wrong way round.

## Install

```bash
uv add --dev torchtyc      # or: pip install torchtyc
```

Extras: `torchtyc[lsp]` for the language server, `[watch]` for watch mode,
`[einops]` for einops-aware hints, `[all]` for everything.

torchtyc must run under the same interpreter as your project, since it imports
your code. By default it finds `.venv/bin/python` next to your `pyproject.toml`.
Override with `--python` or `[tool.torchtyc] python = "..."`.

## Commands

```
torchtyc check <paths>...        Shape-check files or directories
torchtyc trace <file.py::func>   Show the shapes flowing through one function
torchtyc watch <paths>...        Re-check on change
torchtyc lsp                     Language server on stdio
torchtyc lsp --tcp PORT          Language server on a TCP port
torchtyc rules                   List the diagnostic rules
torchtyc version                 Print the version
```

`check` takes `--format full` (default), `concise`, `json`, or `github`.

```
$ torchtyc trace model.py::Linear.forward
Linear.forward
  x       : (..., in_features)
  return -> float32[(..., out_features)]

dimension names are bound to distinct primes starting at 101
```

## Constructing modules

To check a method, torchtyc needs an instance, and to build an instance it needs
constructor arguments. It matches integer parameters of `__init__` against the
dimension names in the method's annotations:

```python
def __init__(self, in_features: int, out_features: int) -> None: ...
def forward(self, x: Float[Tensor, "... in_features"]) -> ...
```

`in_features` is a dimension name, so it receives that dimension's prime.
Parameters with defaults are left alone. A parameter that is neither a
dimension, a known type, nor defaulted produces an `unresolved-arg` warning and
the function is skipped rather than guessed at.

Modules are built inside `torch.device("meta")`, and the initialisers in
`torch.nn.init` are neutralised for the duration, since initial values cannot
affect a shape.

## Annotated attributes

Python does not check variable annotations at runtime, and neither does
jaxtyping, which only reads function signatures. Since torchtyc has a
constructed instance in hand anyway, it checks them too:

```python
# attr.py
class Linear(nn.Module):
    def __init__(self, d_in: int, d_out: int) -> None:
        super().__init__()
        self.W: Float[nn.Parameter, "d_out d_in"] = nn.Parameter(torch.empty((d_in, d_out)))
```

```
attr.py:9:17: error[attribute-mismatch]
  `self.W`: annotated `d_out`, but the traced dimension is `d_in`
    Expected: (d_out, d_in)
    Got:      (d_in, d_out)
    9 | self.W: Float[nn.Parameter, "d_out d_in"] = nn.Parameter(torch.empty((d_in, d_out)))
  try:  Float[nn.Parameter, "d_in d_out"]
  hint: this dimension is `d_in`, so the annotation likely names the wrong axis
```

## Suppressing

```python
y = x.reshape(-1)  # torchtyc: ignore
y = x.reshape(-1)  # torchtyc: ignore[rank-mismatch]
```

A scoped ignore that never matches is itself reported, so suppressions do not
rot.

## Configuration

```toml
[tool.torchtyc]
python = ".venv/bin/python"   # interpreter that imports your code
severity = "warning"          # drop anything below this level
ignore = ["unused-dim"]
exclude = [".venv", "build", "experiments"]
variadic-rank = 2             # how many axes `...` stands for
einops = true
timeout = 60.0
extra-paths = ["stubs"]       # prepended to the worker's PYTHONPATH
allow-effects = false         # let imported code write files, open sockets, and spawn processes
```

## Editors

Any LSP client works. Neovim, without a plugin:

```lua
vim.lsp.config.torchtyc = {
  cmd = { "torchtyc", "lsp" },
  filetypes = { "python" },
  root_markers = { "pyproject.toml", ".git" },
}
vim.lsp.enable("torchtyc")
```

The server publishes lint diagnostics immediately on every change, and traces
after the buffer has been quiet for 0.7s, on open, and on save. It never imports
your code on a keystroke.

Hover over a function to see the traced shapes. Inlay hints show the traced
return next to each signature, and a code lens above it reports the traced
return or the error count. Code actions offer to silence a rule or to adopt
the shape that was actually traced.

## CI

```yaml
- run: uv run torchtyc check src/ --format github
```

`--format github` emits workflow commands, so each finding becomes an inline
annotation on the pull request diff. `check` exits 1 when there are errors, and
2 when torchtyc itself could not do the job: the worker failed, or the paths
matched no python file.

## Rules

| Rule | Level | Meaning |
| --- | --- | --- |
| `shape-mismatch` | error | a traced shape disagrees with its annotation |
| `rank-mismatch` | error | a traced value has a different number of dimensions |
| `dtype-mismatch` | error | a traced dtype is outside the annotated dtype set |
| `dim-inconsistent` | error | one dimension name is bound to two different sizes |
| `attribute-mismatch` | error | an annotated attribute on self holds a different shape |
| `not-a-tensor` | error | an annotated tensor position received a non-tensor |
| `tuple-arity` | error | a tuple return has a different length than annotated |
| `einops-pattern` | error | an einops pattern disagrees with the tensors given to it |
| `trace-error` | error | the function raised while being traced |
| `import-error` | error | the module could not be imported |
| `device-mismatch` | warning | a traced value left the meta device |
| `einops-unknown-axis` | warning | an einops axis matches no input axis or keyword |
| `uninstantiable` | warning | a module's `__init__` could not be called automatically |
| `unresolved-arg` | warning | a parameter has no annotation and no default |
| `unsupported-annotation` | warning | an annotation could not be parsed |
| `local-definition` | info | a target inside a function body cannot be reached after import |
| `anonymous-return` | info | arguments are annotated but the return is not |
| `missing-annotation` | info | a public function has no jaxtyping annotation |
| `unused-dim` | info | a dimension name is used once, so it constrains nothing |
| `suppression-unused` | info | an ignore comment matched no diagnostic |

## Limits

torchtyc imports your module, so module-level side effects run. Keep training
loops behind `if __name__ == "__main__":`.

Classes and functions nested inside other classes, or written under a
module-level `if` or `try`, are checked like any other. One inside a *function
body* is not: after import there is no name to reach it by, so it reports
`local-definition` rather than passing silently. A `if TYPE_CHECKING:` block is
skipped without a diagnostic, because nothing in it exists at runtime.

It runs one concrete trace, not a proof. A function whose control flow depends
on tensor *values* rather than shapes takes whichever branch the primes send it
down. `...` stands for a fixed number of axes, two by default, so code that
behaves differently at other ranks needs `variadic-rank` or a second annotated
wrapper. Every bare `...` in one signature also stands for the *same* axes,
which is narrower than jaxtyping: a loss taking two `"... d"` arguments traces
with one batch shape, because that is what the annotation almost always means.

A dimension is a prime, and a prime does not divide, so code that splits an
axis - `head_dim = d_model // n_heads`, then `view(b, s, n_heads, head_dim)` -
gets a quotient that does not multiply back. A trace that fails this way is run
again on widths that do divide, taken from the numbers the model itself writes
down: a parameter's default, or a literal in the body of the constructor or
the traced method. Multi-head attention passes on the second attempt. A model
that splits by a width written nowhere either of those can see, or by one above
256, still needs `# torchtyc: ignore[trace-error]`.

Runtime checking with `jaxtyping` and `beartype` remains worth having. torchtyc
tells you the shapes are consistent for the sizes it chose; beartype tells you
they were right for the batch you actually ran.

## Security

torchtyc imports your modules in a worker subprocess that runs with your full
privileges.

While it imports and traces your code, a guard blocks filesystem writes,
outbound network calls (including DNS resolution), and process spawning. It
catches accidental effects such as telemetry setup, dataset downloads (including
those shelling out to `git` or `curl`), and checkpoint writes. It is not a
security boundary: Python cannot enforce one in-process, and code written to get
around it can. Set `allow-effects = true` if a model you trust needs to write,
connect, or spawn processes.

Untrusted code still needs a real sandbox or a disposable container. Only enable
`torchtyc lsp` in workspaces you trust.

## Licence

Apache-2.0.
