torchtyc 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- torchtyc/__init__.py +20 -0
- torchtyc/annotations.py +215 -0
- torchtyc/binding.py +479 -0
- torchtyc/cli.py +180 -0
- torchtyc/config.py +134 -0
- torchtyc/diagnostics.py +150 -0
- torchtyc/discovery.py +728 -0
- torchtyc/einops_rules.py +165 -0
- torchtyc/engine.py +409 -0
- torchtyc/formats.py +164 -0
- torchtyc/lsp.py +501 -0
- torchtyc/tracing.py +706 -0
- torchtyc/worker.py +391 -0
- torchtyc-0.1.0.dist-info/METADATA +305 -0
- torchtyc-0.1.0.dist-info/RECORD +18 -0
- torchtyc-0.1.0.dist-info/WHEEL +4 -0
- torchtyc-0.1.0.dist-info/entry_points.txt +2 -0
- torchtyc-0.1.0.dist-info/licenses/LICENSE +21 -0
torchtyc/__init__.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""torchtyc: static array shape checking for PyTorch, powered by meta tensors."""
|
|
2
|
+
|
|
3
|
+
from .config import Config, load
|
|
4
|
+
from .diagnostics import RULES, Diagnostic, Rule, Severity
|
|
5
|
+
from .engine import Report, check_paths, collect_files
|
|
6
|
+
|
|
7
|
+
__version__ = "0.1.0"
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"RULES",
|
|
11
|
+
"Config",
|
|
12
|
+
"Diagnostic",
|
|
13
|
+
"Report",
|
|
14
|
+
"Rule",
|
|
15
|
+
"Severity",
|
|
16
|
+
"__version__",
|
|
17
|
+
"check_paths",
|
|
18
|
+
"collect_files",
|
|
19
|
+
"load",
|
|
20
|
+
]
|
torchtyc/annotations.py
ADDED
|
@@ -0,0 +1,215 @@
|
|
|
1
|
+
"""Parsing of jaxtyping annotations out of source text.
|
|
2
|
+
|
|
3
|
+
Everything here works on `ast` nodes and strings only. It never imports torch or
|
|
4
|
+
evaluates the annotation, so the discovery pass can run in the editor's process
|
|
5
|
+
on a file that does not even import cleanly.
|
|
6
|
+
|
|
7
|
+
The grammar follows jaxtyping's dim strings:
|
|
8
|
+
|
|
9
|
+
"batch seq d_model" named dimensions
|
|
10
|
+
"3 seq" a fixed size
|
|
11
|
+
"..." an anonymous run of dimensions
|
|
12
|
+
"*batch" a named run of dimensions
|
|
13
|
+
"#channels" a broadcastable dimension
|
|
14
|
+
"_" one anonymous dimension
|
|
15
|
+
"d_in+d_out" a symbolic expression over already-bound names
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import ast
|
|
21
|
+
import re
|
|
22
|
+
from dataclasses import dataclass
|
|
23
|
+
|
|
24
|
+
# jaxtyping's dtype classes. The values are resolved to concrete torch dtypes in
|
|
25
|
+
# the worker; here they are only names, so this table stays importable anywhere.
|
|
26
|
+
DTYPE_NAMES = frozenset(
|
|
27
|
+
{
|
|
28
|
+
"Shaped",
|
|
29
|
+
"Num",
|
|
30
|
+
"Inexact",
|
|
31
|
+
"Real",
|
|
32
|
+
"Float",
|
|
33
|
+
"Complex",
|
|
34
|
+
"Integer",
|
|
35
|
+
"Int",
|
|
36
|
+
"UInt",
|
|
37
|
+
"Bool",
|
|
38
|
+
"Key",
|
|
39
|
+
"Float16",
|
|
40
|
+
"Float32",
|
|
41
|
+
"Float64",
|
|
42
|
+
"BFloat16",
|
|
43
|
+
"Float8e4m3fn",
|
|
44
|
+
"Float8e5m2",
|
|
45
|
+
"Complex64",
|
|
46
|
+
"Complex128",
|
|
47
|
+
"Int4",
|
|
48
|
+
"Int8",
|
|
49
|
+
"Int16",
|
|
50
|
+
"Int32",
|
|
51
|
+
"Int64",
|
|
52
|
+
"UInt4",
|
|
53
|
+
"UInt8",
|
|
54
|
+
"UInt16",
|
|
55
|
+
"UInt32",
|
|
56
|
+
"UInt64",
|
|
57
|
+
}
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
_IDENT = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
|
61
|
+
_SYMBOLIC = re.compile(r"[+\-*/()]")
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
@dataclass(frozen=True)
|
|
65
|
+
class Dim:
|
|
66
|
+
"""One entry in a dim string."""
|
|
67
|
+
|
|
68
|
+
kind: str # "named" | "fixed" | "anonymous" | "variadic" | "symbolic"
|
|
69
|
+
name: str | None = None
|
|
70
|
+
size: int | None = None
|
|
71
|
+
expr: str | None = None
|
|
72
|
+
broadcastable: bool = False
|
|
73
|
+
|
|
74
|
+
def __str__(self) -> str:
|
|
75
|
+
prefix = "#" if self.broadcastable else ""
|
|
76
|
+
if self.kind == "fixed":
|
|
77
|
+
return str(self.size)
|
|
78
|
+
if self.kind == "anonymous":
|
|
79
|
+
return "_"
|
|
80
|
+
if self.kind == "variadic":
|
|
81
|
+
return "..." if self.name is None else f"*{self.name}"
|
|
82
|
+
if self.kind == "symbolic":
|
|
83
|
+
return f"{prefix}{self.expr}"
|
|
84
|
+
return f"{prefix}{self.name}"
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
@dataclass(frozen=True)
|
|
88
|
+
class ArraySpec:
|
|
89
|
+
"""A single `Dtype[ArrayType, "dims"]` annotation."""
|
|
90
|
+
|
|
91
|
+
dtype: str
|
|
92
|
+
array_type: str
|
|
93
|
+
dims: tuple[Dim, ...]
|
|
94
|
+
raw: str
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def named_dims(self) -> tuple[str, ...]:
|
|
98
|
+
return tuple(d.name for d in self.dims if d.name is not None)
|
|
99
|
+
|
|
100
|
+
def __str__(self) -> str:
|
|
101
|
+
return f'{self.dtype}[{self.array_type}, "{" ".join(str(d) for d in self.dims)}"]'
|
|
102
|
+
|
|
103
|
+
def shape_str(self) -> str:
|
|
104
|
+
return "(" + ", ".join(str(d) for d in self.dims) + ")"
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
@dataclass(frozen=True)
|
|
108
|
+
class TupleSpec:
|
|
109
|
+
"""A `tuple[...]` of array specs in a return position."""
|
|
110
|
+
|
|
111
|
+
items: tuple[Spec, ...]
|
|
112
|
+
|
|
113
|
+
def __str__(self) -> str:
|
|
114
|
+
return "tuple[" + ", ".join(str(i) for i in self.items) + "]"
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
@dataclass(frozen=True)
|
|
118
|
+
class OpaqueSpec:
|
|
119
|
+
"""An annotation torchtyc does not model, kept so callers can report it."""
|
|
120
|
+
|
|
121
|
+
raw: str
|
|
122
|
+
|
|
123
|
+
def __str__(self) -> str:
|
|
124
|
+
return self.raw
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
Spec = ArraySpec | TupleSpec | OpaqueSpec
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
class AnnotationError(ValueError):
|
|
131
|
+
"""The annotation looked like jaxtyping but could not be parsed."""
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def parse_dim_string(text: str) -> tuple[Dim, ...]:
|
|
135
|
+
"""Split a jaxtyping dim string into dims, left to right."""
|
|
136
|
+
dims: list[Dim] = []
|
|
137
|
+
for token in text.split():
|
|
138
|
+
dims.append(_parse_dim_token(token))
|
|
139
|
+
return tuple(dims)
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def _parse_dim_token(token: str) -> Dim:
|
|
143
|
+
broadcastable = token.startswith("#")
|
|
144
|
+
if broadcastable:
|
|
145
|
+
token = token[1:]
|
|
146
|
+
if token == "...":
|
|
147
|
+
return Dim("variadic")
|
|
148
|
+
if token.startswith("*"):
|
|
149
|
+
rest = token[1:]
|
|
150
|
+
if rest and not _IDENT.match(rest):
|
|
151
|
+
raise AnnotationError(f"bad variadic dimension {token!r}")
|
|
152
|
+
return Dim("variadic", name=rest or None, broadcastable=broadcastable)
|
|
153
|
+
if token == "_":
|
|
154
|
+
return Dim("anonymous", broadcastable=broadcastable)
|
|
155
|
+
if token.isdigit():
|
|
156
|
+
return Dim("fixed", size=int(token), broadcastable=broadcastable)
|
|
157
|
+
if _SYMBOLIC.search(token):
|
|
158
|
+
return Dim("symbolic", expr=token, broadcastable=broadcastable)
|
|
159
|
+
if _IDENT.match(token):
|
|
160
|
+
return Dim("named", name=token, broadcastable=broadcastable)
|
|
161
|
+
raise AnnotationError(f"bad dimension {token!r}")
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def _dotted_name(node: ast.expr) -> str | None:
|
|
165
|
+
"""Render `Tensor` or `torch.Tensor` or `nn.Parameter` back to a string."""
|
|
166
|
+
if isinstance(node, ast.Name):
|
|
167
|
+
return node.id
|
|
168
|
+
if isinstance(node, ast.Attribute):
|
|
169
|
+
base = _dotted_name(node.value)
|
|
170
|
+
return None if base is None else f"{base}.{node.attr}"
|
|
171
|
+
return None
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
def parse_annotation(node: ast.expr | None) -> Spec | None:
|
|
175
|
+
"""Turn an annotation node into a spec, or None if there is no annotation.
|
|
176
|
+
|
|
177
|
+
Anything that is not recognisably jaxtyping comes back as an OpaqueSpec so
|
|
178
|
+
the caller can decide whether to warn or stay quiet.
|
|
179
|
+
"""
|
|
180
|
+
if node is None:
|
|
181
|
+
return None
|
|
182
|
+
|
|
183
|
+
if isinstance(node, ast.Subscript):
|
|
184
|
+
head = _dotted_name(node.value)
|
|
185
|
+
tail = head.rsplit(".", 1)[-1] if head else None
|
|
186
|
+
|
|
187
|
+
if tail in DTYPE_NAMES:
|
|
188
|
+
return _parse_array(node, tail)
|
|
189
|
+
|
|
190
|
+
if tail in ("tuple", "Tuple"):
|
|
191
|
+
items = node.slice.elts if isinstance(node.slice, ast.Tuple) else [node.slice]
|
|
192
|
+
parsed = tuple(parse_annotation(item) or OpaqueSpec("?") for item in items)
|
|
193
|
+
if any(isinstance(p, ArraySpec) for p in parsed):
|
|
194
|
+
return TupleSpec(parsed)
|
|
195
|
+
|
|
196
|
+
return OpaqueSpec(ast.unparse(node))
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
def _parse_array(node: ast.Subscript, dtype: str) -> Spec:
|
|
200
|
+
slice_node = node.slice
|
|
201
|
+
if not isinstance(slice_node, ast.Tuple) or len(slice_node.elts) != 2:
|
|
202
|
+
raise AnnotationError(f'expected {dtype}[ArrayType, "dims"]')
|
|
203
|
+
|
|
204
|
+
array_node, dims_node = slice_node.elts
|
|
205
|
+
array_type = _dotted_name(array_node) or ast.unparse(array_node)
|
|
206
|
+
|
|
207
|
+
if not isinstance(dims_node, ast.Constant) or not isinstance(dims_node.value, str):
|
|
208
|
+
raise AnnotationError("the second argument must be a literal dim string")
|
|
209
|
+
|
|
210
|
+
return ArraySpec(
|
|
211
|
+
dtype=dtype,
|
|
212
|
+
array_type=array_type,
|
|
213
|
+
dims=parse_dim_string(dims_node.value),
|
|
214
|
+
raw=ast.unparse(node),
|
|
215
|
+
)
|