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 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
+ ]
@@ -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
+ )