Stub Library Author Reference
This page documents the shape_extensions features used to describe tensor
library APIs. Application authors normally need only the
tensor shape API reference.
How shape DSL rules run
A shape rule is a function decorated with @type_shape_dsl_function. A public
stub calls that function in its return annotation:
import shape_extensions.dsl as dsl
from shape_extensions import IntTuple, type_shape_dsl_function
from torch import Tensor
@type_shape_dsl_function
def repeat_shape(shape: IntTuple, repeats: IntTuple) -> IntTuple:
if len(repeats) < len(shape):
return dsl.Invalid("repeat dimensions cannot be shorter than the input rank")
extra = len(repeats) - len(shape)
return dsl.IntTuple(
repeats[i] if i < extra else shape[i - extra] * repeats[i]
for i in range(len(repeats))
)
def repeat[Shape: IntTuple, Repeats: IntTuple](
self: Tensor[Shape], *sizes: *Repeats
) -> Tensor[repeat_shape(Shape, Repeats)]: ...
Calls to type-level DSL functions are valid only in return annotations. When
Pyrefly solves a call to repeat, it binds Shape and Repeats from the
arguments, immediately evaluates repeat_shape(Shape, Repeats), and uses the
result as the returned Tensor shape. The DSL function is not called at
runtime; its decorator is a runtime no-op.
Use a plain generic signature whenever it can express the relationship. Add a
DSL rule only when the return shape requires computation, as in Torch
reshape, NumPy stack, or JAX conv_general_dilated.
Annotation constructs
These constructs connect ordinary Python arguments to shape types. They are
intended for stub libraries; application code should normally use Tensor,
Int, IntVar, IntTuple, and Elements as described in the
API reference.
| Construct | Stub-author use |
|---|---|
IntTuples | Carries an ordered collection of shapes, such as the operands of Torch einsum or JAX concatenate. |
Flag[T] | Preserves a runtime option as a literal-like type argument so a rule can inspect it, such as axis in NumPy sum or dim and keepdim in Torch sum. |
Index with index_shape | Captures a complete indexing expression for shape-aware Tensor.__getitem__, including tuples, slices, and advanced indices. |
IntListLiteral[S] | Captures a direct integer list literal as shape S; it is normally exposed through IntTupleOrList[S]. |
IntTupleOrList[S] | Accepts either a tuple or direct list of integers while retaining its values, as in torch.split(x, [2, 3]). |
RegularNestedList[S, D] | Infers shape S and scalar domain D from a regular nested literal passed to numpy.array, jax.numpy.array, or torch.tensor. |
NamedInts with CaptureNamedInts | Captures named integer **kwargs, such as copies=4 in an einops repeat call, for use by an einops shape rule. |
ProxyMethod["name"] | Gives a forwarding method the signature of another method; Torch uses ProxyMethod["forward"] for nn.Module.__call__. |
IntListLiteral and RegularNestedList apply contextual typing only to direct
list literals. Existing lists retain their ordinary Python types so these
markers do not claim shape information that is unavailable at the call site.
MapIntTuples maps in both directions
MapIntTuples[lambda S: Array[S], Shapes] applies an array type constructor to
every shape in Shapes. In a parameter annotation it runs in reverse: NumPy
stack uses the pattern to infer Shapes from a sequence of differently shaped
arrays.
def stack[Shapes: IntTuples, Axis: Flag[int]](
arrays: MapIntTuples[lambda S: ndarray[S], Shapes],
axis: Axis = 0,
) -> ndarray[stack_shape(Shapes, Axis)]: ...
In a return annotation it runs forward: Torch meshgrid maps each shape
computed by meshgrid_shapes to a corresponding Tensor result.
def meshgrid[Shapes: IntTuples, Indexing: Flag[str | None]](
*tensors: Unpack[MapIntTuples[lambda S: Tensor[S], Shapes]],
indexing: Indexing = None,
) -> MapIntTuples[
lambda S: Tensor[S], meshgrid_shapes(Shapes, Indexing)
]: ...
Other library hooks
static_jaxtypingenables Pyrefly's supported jaxtyping syntax for a declaration; see Jaxtyping compatibility.defines_assert_shapemarks a library-defined equivalent ofassert_shapeso Pyrefly checks calls to it the same way.assert_shapechecks inferred and runtime shapes in the NumPy, Torch, and JAX stub test suites.assert_raiseslets the same test express an expected static error and runtime exception.SymbolicArithExpris the runtime representation used when an evaluated annotation contains symbolic dimension arithmetic; do not use it in stubs.
DSL language
DSL parameters use Int for one dimension, IntTuple for one shape,
IntTuples for a collection of shapes, and NamedInts for captured named
integers. Runtime configuration reaches the helper as ordinary int, bool,
str, integer-tuple, or None values after a public signature captures it with
Flag[...]. A helper returns Int, IntTuple, or IntTuples, or
dsl.Invalid(...) when the call is ill-formed.
The function body is a deliberately small, immutable Python subset:
- an optional docstring, single-assignment local names,
if/else, and explicitreturnstatements; - integer arithmetic (
+,-,*,//,%), comparisons, boolean operations, and conditional expressions; - shape indexing and slicing, plus
len,range,tuple,zip,any, and tuplecountandindex; - bounded comprehensions and generator expressions over those values.
Mutation, loops, exceptions, and arbitrary Python calls are not supported. A DSL helper may call another decorated helper only by returning that call directly. Pyrefly validates the definition itself, so unsupported syntax is reported in the stub rather than silently becoming a runtime dependency.
Constructors and control primitives
| Primitive | Purpose |
|---|---|
dsl.IntTuple(values) | Builds one shape from a fixed tuple or bounded generator, as Torch repeat_shape does. |
dsl.IntTuples(values) | Builds a collection of shapes, as Torch meshgrid_shapes does. |
dsl.concat(left, right) | Concatenates shape fragments; Torch unsqueeze_shape uses it to append a dimension. |
dsl.prod(shape) / dsl.sum(shape) | Reduces dimensions to one symbolic integer; reshape rules use prod, while Torch cat_shape uses sum. |
dsl.is_concrete_int(value) | Narrows an Int only when its value is known, for checks such as Torch select bounds. |
dsl.is_int_value(value) | Narrows the integer arm of a flag union, such as Torch squeeze's `int |
dsl.Invalid(message) | Rejects an invalid library call with a shape diagnostic, such as an out-of-range Torch dimension. |
Use dsl.Int.gradual(), dsl.IntTuple.gradual(), or
dsl.IntTuples.gradual() when the call is valid but its result cannot be
determined statically. A gradual result preserves whatever outer structure is
known; dsl.Invalid(...) instead reports an error.
Reusable shape operations
| Operation | Purpose |
|---|---|
broadcast(left, right) | Computes ordinary two-input broadcasting, used by elementwise Torch, NumPy, and JAX operators. |
gufunc_broadcast(spec, shapes) | Applies a generalized-ufunc core-dimension signature, such as "(m,n),(n,p)->(m,p)" for NumPy or JAX matrix multiplication. |
dsl.einsum(spec, shapes) | Computes shapes for explicit single-letter einsum equations, as in torch.einsum("ij,jk->ik", ...). |
dsl.einops_einsum(spec, shapes) | Computes shapes for einops named-axis einsum equations. |
dsl.rearrange(spec, shape, axes) | Applies an einops rearrangement pattern and optional NamedInts axis lengths. |
dsl.reduce(spec, shape, axes) | Applies an einops reduction pattern and optional named axis lengths. |
dsl.repeat(spec, shape, axes) | Applies an einops repeat pattern and optional named axis lengths. |
broadcast and gufunc_broadcast are complete decorated rules exported from
shape_extensions; call them from return annotations or return them directly
from another DSL helper. The other operations are primitives in
shape_extensions.dsl and are available only inside DSL definitions.