Skip to main content

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.

ConstructStub-author use
IntTuplesCarries 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_shapeCaptures 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 CaptureNamedIntsCaptures 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_jaxtyping enables Pyrefly's supported jaxtyping syntax for a declaration; see Jaxtyping compatibility.
  • defines_assert_shape marks a library-defined equivalent of assert_shape so Pyrefly checks calls to it the same way.
  • assert_shape checks inferred and runtime shapes in the NumPy, Torch, and JAX stub test suites. assert_raises lets the same test express an expected static error and runtime exception.
  • SymbolicArithExpr is 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 explicit return statements;
  • integer arithmetic (+, -, *, //, %), comparisons, boolean operations, and conditional expressions;
  • shape indexing and slicing, plus len, range, tuple, zip, any, and tuple count and index;
  • 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​

PrimitivePurpose
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​

OperationPurpose
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.