Skip to content

Shaped annotations that other type checkers can read [rfc] (#5027) - #5027

Draft
stroxler wants to merge 1 commit into
mainfrom
export-D122076108
Draft

stroxler wants to merge 1 commit into
mainfrom
export-D122076108

Conversation

@stroxler

@stroxler stroxler commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor

Summary:

Prototype, for gathering feedback. Not intended to land as is.

Pyrefly's native shape annotations put Python expressions where the typing spec
allows only types, so pyright, mypy and ty reject them. A library that wants to
track shapes therefore cannot publish annotations that its users can check with
other tools. This prototype carries the shape as string metadata instead:

from shape_extensions import Shaped, shape_vars

shape_vars("M, K, N")
def matmul(
    a: Shaped[np.ndarray, "[M, K]"],
    b: Shaped[np.ndarray, "[K, N]"],
) -> Shaped[np.ndarray, "[M, N]"]: ...

shape_vars("M, N")
def transpose(
    x: np.ndarray[Shaped[tuple[int, int], "[M, N]"], np.dtype[np.float64]],
) -> np.ndarray[Shaped[tuple[int, int], "[N, M]"], np.dtype[np.float64]]: ...

shape_vars("N")
def size(x: Shaped[np.ndarray, "[N]"]) -> Shaped[int, "N"]: ...

shape_vars("Dim, Hidden")
class Encoder:
    weight: Shaped[np.ndarray, "[Dim, Hidden]"]

shape_vars("Dim, Hidden")
def output_size(encoder: Shaped[Encoder, "Dim, Hidden"]) -> Shaped[int, "Hidden"]: ...

To Pyrefly these are shape-checked signatures. To every other checker Shaped is
typing.Annotated, so they see ndarray, ndarray[tuple[int, int], dtype],
int and Encoder.

Surface

  • shape_extensions.Shaped is from typing import Annotated as Shaped. It has to
    be an import alias, because mypy rejects Shaped = Annotated as not a valid
    type. Pyrefly still recognizes it as Shaped.
  • shape_vars("M, N, *S") declares the dimensions a function or class may use.
    A bare name is a single dimension (an IntVar) and *S is a variadic run of
    dimensions (an IntTuple). No bounds are written. It is the counterpart of
    static_jaxtyping for jaxtyping, and the two stay separate.
  • On a class, the declared dimensions become type parameters that follow the
    class's own type parameters, so Encoder above is generic in Dim and
    Hidden. Only Pyrefly sees these parameters. They default to gradual
    dimensions, so adding shape_vars to a published class breaks no consumer:
    for a generic shape_vars("A, B") class GEncoder[T], a consumer's
    GEncoder[Input] still checks under Pyrefly, and its dimensions are unknown.
    A library that wants Pyrefly users to spell the dimensions out can write
    shape_vars("A, B", required=True), and then GEncoder[Input] is an arity
    error. Either way, other checkers see only GEncoder[T].
  • Shaped[T, "..."] means something only inside a shape_vars scope. There, the
    string is parsed with the native shape syntax: a list [M, N] is an integer
    tuple, and a bare name or expression is a dimension. Names resolve against the
    declaration first, then ordinary scopes. Outside a scope Shaped[T, "..."] is
    plain Annotated, meaning T, exactly as other checkers read it. A literal-only
    shape such as "[3, 4]" also needs a scope, which can be empty: shape_vars("").

Semantics: two modes

Pyrefly decides the meaning of Shaped[T, "args"] from what T resolves to.

Append mode (any generic class). The string's comma-separated arguments are
appended after T's explicit type arguments, and the result is an ordinary
specialization, checked like any other:

  • Shaped[np.ndarray, "[M, N]"] is np.ndarray[[M, N]]. With the Pyrefly numpy
    stubs this is ndarray[IntTuple[M, N]]. With real numpy the list fills
    ndarray's tuple-bounded shape parameter, so the shape is still tracked.
  • Shaped[Encoder, "Dim, Hidden"] is Encoder[Dim, Hidden].
  • For a generic class with its own type parameters, Shaped[GEncoder[Input], "A, B"] is GEncoder[Input, A, B].
  • A wrong number or kind of argument is an ordinary type-argument error, such as
    "Expected 2 type arguments for Encoder, got 4".

A class with no type parameters, such as real torch's Tensor, ignores the shape
without an error. This lets a library annotated against Pyrefly's torch stubs
still check against real torch.

Replace mode (int, a tuple, or Any). Here T is replaced by the symbolic
type that sits on top of it:

  • Shaped[int, "N"] is Int[N], in any position. A literal argument solves N
    exactly: with size above, zeros(3) has shape [3]. To other checkers it is
    int, and Int[N] is assignable to int.
  • Shaped[tuple[int, int], "[M, N]"] is IntTuple[M, N], anywhere in an
    annotation: a type argument, a union, a parameter or a return type. To other
    checkers it is the tuple. A function can therefore take or return a shape,
    as in def shape_of(x: Shaped[ndarray, "[M, N]"]) -> Shaped[tuple[int, int], "[M, N]"], and a literal tuple such as (3, 4) is accepted where
    IntTuple[M, N] is expected. Inside an array's shape argument, this is the
    spelling for an array whose dtype matters, such as
    ndarray[Shaped[tuple[int], "[N]"], np.dtype[np.float64]], because other
    checkers then keep numpy's own rank and dtype. The tuple must agree with the
    shape: every rank the shape allows must be one the tuple allows, and every
    literal the tuple states must appear at the same position in the shape.
    Bare tuple, tuple[int, ...] and explicit Any state no rank, so they agree
    with any shape. A tuple that disagrees is reported and means the tuple.

Replace mode takes precedence over append mode. tuple is generic, but
appending a shape to it would give a meaningless tuple[int, int, [M, N]].

Malformed annotations. An unparsable string or a failed replacement is
reported, and the annotation then means T, as it does to other checkers. In
append mode, errors are ordinary type-argument errors and the result is what
Pyrefly infers for the bad specialization, which may differ from T.

Cross-checker tests

compatibility_tests checks that a library annotated this way is still readable
by tools that know nothing about shapes. It pairs an annotated library with a
downstream consumer that never imports shape_extensions, and runs Pyrefly with
the stubs, Pyrefly with real numpy, mypy, pyright and ty. The consumer's
assert_type calls fail both if the metadata leaks into the type and if the type
degrades to Any. tensor-shapes/run_tests.py runs the suite.

Open Questions / Potential Improvements or Changes

  • Decorators new decorators in stubs is not currently permitted by the typing spec,
    so other type checkers may complain about @shape_vars. In addition, the
    spec generally wants decorators to not be mixed with overloads, but in this
    case the decorator is sugar for type params so it needs to be per-overload.
    Both problems could potentially be worked around by instead enabling some
    kind of comment directive for Pyrefly to desugar in the same way, but that
    would be a more challenging
  • Only annotations, assert_type and type parameter bounds read the string.
    In a cast or a type alias, and inside a class base's type arguments, Shaped
    silently drops its shape. A bug test pins this, it might be possible to improve
    (cast in particular seems like a significant gap)

Limits of the current approach

  • A bare generic base shifts the shape. Append mode appends after the
    arguments written, so Shaped[GEncoder, "A, B"] binds A to GEncoder's own
    parameter T. The class's own arguments must be written out, as in
    Shaped[GEncoder[Any], "A, B"]. Whether Shaped on a shape_vars class
    should instead fill its declared dimensions is an open question.
  • Appending follows ordinary subscripting. Shaped[C[args], "more"] means
    C[args, more], so a base that cannot be subscripted cannot take a shape. A
    non-generic alias such as type Array = np.ndarray reports "not
    subscriptable", exactly as Array[int] does.
  • Subclasses cannot carry dimensions portably. class Sub(Encoder[3, 4])
    runs, because shape_vars makes the class subscriptable at runtime, but only
    Pyrefly accepts it. class Sub(Shaped[Encoder, "3, 4"]) also runs, and mypy
    and ty accept it, but pyright rejects it ("Argument to class must be a base
    class") and Pyrefly reports "Invalid base class: Annotated".
  • required=True cannot follow a defaulted type parameter. A required
    dimension without a default cannot follow class C[T = int], as for any type
    parameter.
  • A class without type parameters silently ignores the shape, which is
    needed for real torch but hides a forgotten shape_vars on a user class.
  • Variadic shapes against real numpy. A variadic tuple is rejected because
    IntTuple[*S] is not assignable to numpy's tuple[int, ...] bound. I'm not sure
    whether this actually matters, since tuple[Any, ...] == tuple is actually the
    variadic form for other type checkers.
  • Pyright evaluates Annotated string metadata lazily as a forward
    reference. It accepts undefined bare names, lists and arithmetic, but not
    calls, so a shape function named in a string must be imported.
  • Ruff recognizes Annotated only under that name, so it reports the names
    in a Shaped string as undefined (F821). For now, code that uses Shaped
    disables F821, as compatibility_tests does.

Implementation notes

Inside a shape_vars scope the binder parses the shape string and binds its
names. The solver picks the mode from the resolved base in alt/shaped.rs. The
shape hooks live in alt/shaped.rs, alt/shape_extension.rs,
alt/shape_declarations.rs and binding/shape_type.rs. Core also gains a small
amount of plumbing: a binding for declared dimensions, and a name-resolution
override for names in a shape string.

Differential Revision: D122076108

@meta-cla meta-cla Bot added the cla signed label Sep 28, 2026
@meta-codesync

meta-codesync Bot commented Sep 28, 2026

Copy link
Copy Markdown
Contributor

@stroxler has exported this pull request. If you are a Meta employee, you can view the originating Diff in D122076108.

@github-actions github-actions Bot added size/xl and removed size/xl labels Sep 28, 2026
@stroxler stroxler changed the title Prototype: Shaped annotations that other type checkers can read RFC: Shaped annotations that other type checkers can read Sep 28, 2026
Summary:
Pull Request resolved: #5027

Prototype, for gathering feedback. Not intended to land as is.

Pyrefly's native shape annotations put Python expressions where the typing spec
allows only types, so pyright, mypy and ty reject them. A library that wants to
track shapes therefore cannot publish annotations that its users can check with
other tools. This prototype carries the shape as string metadata instead:

    from shape_extensions import Shaped, shape_vars

    shape_vars("M, K, N")
    def matmul(
        a: Shaped[np.ndarray, "[M, K]"],
        b: Shaped[np.ndarray, "[K, N]"],
    ) -> Shaped[np.ndarray, "[M, N]"]: ...

    shape_vars("M, N")
    def transpose(
        x: np.ndarray[Shaped[tuple[int, int], "[M, N]"], np.dtype[np.float64]],
    ) -> np.ndarray[Shaped[tuple[int, int], "[N, M]"], np.dtype[np.float64]]: ...

    shape_vars("N")
    def size(x: Shaped[np.ndarray, "[N]"]) -> Shaped[int, "N"]: ...

    shape_vars("Dim, Hidden")
    class Encoder:
        weight: Shaped[np.ndarray, "[Dim, Hidden]"]

    shape_vars("Dim, Hidden")
    def output_size(encoder: Shaped[Encoder, "Dim, Hidden"]) -> Shaped[int, "Hidden"]: ...

To Pyrefly these are shape-checked signatures. To every other checker `Shaped` is
`typing.Annotated`, so they see `ndarray`, `ndarray[tuple[int, int], dtype]`,
`int` and `Encoder`.

## Surface

- `shape_extensions.Shaped` is `from typing import Annotated as Shaped`. It has to
  be an import alias, because mypy rejects `Shaped = Annotated` as not a valid
  type. Pyrefly still recognizes it as `Shaped`.
- `shape_vars("M, N, *S")` declares the dimensions a function or class may use.
  A bare name is a single dimension (an `IntVar`) and `*S` is a variadic run of
  dimensions (an `IntTuple`). No bounds are written. It is the counterpart of
  `static_jaxtyping` for jaxtyping, and the two stay separate.
- On a class, the declared dimensions become type parameters that follow the
  class's own type parameters, so `Encoder` above is generic in `Dim` and
  `Hidden`. Only Pyrefly sees these parameters. They default to gradual
  dimensions, so adding `shape_vars` to a published class breaks no consumer:
  for a generic `shape_vars("A, B") class GEncoder[T]`, a consumer's
  `GEncoder[Input]` still checks under Pyrefly, and its dimensions are unknown.
  A library that wants Pyrefly users to spell the dimensions out can write
  `shape_vars("A, B", required=True)`, and then `GEncoder[Input]` is an arity
  error. Either way, other checkers see only `GEncoder[T]`.
- `Shaped[T, "..."]` means something only inside a `shape_vars` scope. There, the
  string is parsed with the native shape syntax: a list `[M, N]` is an integer
  tuple, and a bare name or expression is a dimension. Names resolve against the
  declaration first, then ordinary scopes. Outside a scope `Shaped[T, "..."]` is
  plain `Annotated`, meaning `T`, exactly as other checkers read it. A literal-only
  shape such as `"[3, 4]"` also needs a scope, which can be empty: `shape_vars("")`.

## Semantics: two modes

Pyrefly decides the meaning of `Shaped[T, "args"]` from what `T` resolves to.

**Append mode (any generic class).** The string's comma-separated arguments are
appended after `T`'s explicit type arguments, and the result is an ordinary
specialization, checked like any other:

- `Shaped[np.ndarray, "[M, N]"]` is `np.ndarray[[M, N]]`. With the Pyrefly numpy
  stubs this is `ndarray[IntTuple[M, N]]`. With real numpy the list fills
  `ndarray`'s `tuple`-bounded shape parameter, so the shape is still tracked.
- `Shaped[Encoder, "Dim, Hidden"]` is `Encoder[Dim, Hidden]`.
- For a generic class with its own type parameters, `Shaped[GEncoder[Input],
  "A, B"]` is `GEncoder[Input, A, B]`.
- A wrong number or kind of argument is an ordinary type-argument error, such as
  "Expected 2 type arguments for `Encoder`, got 4".

A class with no type parameters, such as real torch's `Tensor`, ignores the shape
without an error. This lets a library annotated against Pyrefly's torch stubs
still check against real torch.

**Replace mode (`int`, a tuple, or `Any`).** Here `T` is replaced by the symbolic
type that sits on top of it:

- `Shaped[int, "N"]` is `Int[N]`, in any position. A literal argument solves `N`
  exactly: with `size` above, `zeros(3)` has shape `[3]`. To other checkers it is
  `int`, and `Int[N]` is assignable to `int`.
- `Shaped[tuple[int, int], "[M, N]"]` is `IntTuple[M, N]`, anywhere in an
  annotation: a type argument, a union, a parameter or a return type. To other
  checkers it is the tuple. A function can therefore take or return a shape,
  as in `def shape_of(x: Shaped[ndarray, "[M, N]"]) -> Shaped[tuple[int, int],
  "[M, N]"]`, and a literal tuple such as `(3, 4)` is accepted where
  `IntTuple[M, N]` is expected. Inside an array's shape argument, this is the
  spelling for an array whose dtype matters, such as
  `ndarray[Shaped[tuple[int], "[N]"], np.dtype[np.float64]]`, because other
  checkers then keep numpy's own rank and dtype. The tuple must agree with the
  shape: every rank the shape allows must be one the tuple allows, and every
  literal the tuple states must appear at the same position in the shape.
  Bare `tuple`, `tuple[int, ...]` and explicit `Any` state no rank, so they agree
  with any shape. A tuple that disagrees is reported and means the tuple.

Replace mode takes precedence over append mode. `tuple` is generic, but
appending a shape to it would give a meaningless `tuple[int, int, [M, N]]`.

**Malformed annotations.** An unparsable string or a failed replacement is
reported, and the annotation then means `T`, as it does to other checkers. In
append mode, errors are ordinary type-argument errors and the result is what
Pyrefly infers for the bad specialization, which may differ from `T`.

## Cross-checker tests

`compatibility_tests` checks that a library annotated this way is still readable
by tools that know nothing about shapes. It pairs an annotated library with a
downstream consumer that never imports `shape_extensions`, and runs Pyrefly with
the stubs, Pyrefly with real numpy, mypy, pyright and ty. The consumer's
`assert_type` calls fail both if the metadata leaks into the type and if the type
degrades to `Any`. `tensor-shapes/run_tests.py` runs the suite.

## Open questions and known limits

- **A bare generic base shifts the shape.** Append mode appends after the
  arguments written, so `Shaped[GEncoder, "A, B"]` binds `A` to `GEncoder`'s own
  parameter `T`. The class's own arguments must be written out, as in
  `Shaped[GEncoder[Any], "A, B"]`. Whether `Shaped` on a `shape_vars` class
  should instead fill its declared dimensions is an open question.
- **Appending follows ordinary subscripting.** `Shaped[C[args], "more"]` means
  `C[args, more]`, so a base that cannot be subscripted cannot take a shape. A
  non-generic alias such as `type Array = np.ndarray` reports "not
  subscriptable", exactly as `Array[int]` does.
- **Subclasses cannot carry dimensions portably.** `class Sub(Encoder[3, 4])`
  runs, because `shape_vars` makes the class subscriptable at runtime, but only
  Pyrefly accepts it. `class Sub(Shaped[Encoder, "3, 4"])` also runs, and mypy
  and ty accept it, but pyright rejects it ("Argument to class must be a base
  class") and Pyrefly reports "Invalid base class: `Annotated`".
- **`required=True` cannot follow a defaulted type parameter.** A required
  dimension without a default cannot follow `class C[T = int]`, as for any type
  parameter.
- **Only annotations, `assert_type` and type parameter bounds read the string.**
  In a `cast` or a type alias, and inside a class base's type arguments, `Shaped`
  silently drops its shape. A `bug` test pins this.
- **A class without type parameters silently ignores the shape**, which is
  needed for real torch but hides a forgotten `shape_vars` on a user class.
- **Variadic shapes against real numpy.** A variadic tuple is rejected because
  `IntTuple[*S]` is not assignable to numpy's `tuple[int, ...]` bound. This
  subtyping gap would need closing before variadic shapes could go upstream.
- **Pyright** evaluates `Annotated` string metadata lazily as a forward
  reference. It accepts undefined bare names, lists and arithmetic, but not
  calls, so a shape function named in a string must be imported.
- **Ruff** recognizes `Annotated` only under that name, so it reports the names
  in a `Shaped` string as undefined (F821). For now, code that uses `Shaped`
  disables F821, as `compatibility_tests` does.

## Implementation notes

Inside a `shape_vars` scope the binder parses the shape string and binds its
names. The solver picks the mode from the resolved base in `alt/shaped.rs`. The
shape hooks live in `alt/shaped.rs`, `alt/shape_extension.rs`,
`alt/shape_declarations.rs` and `binding/shape_type.rs`. Core also gains a small
amount of plumbing: a binding for declared dimensions, and a name-resolution
override for names in a shape string.

Differential Revision: D122076108
@stroxler
stroxler marked this pull request as draft September 28, 2026 14:26
@meta-codesync meta-codesync Bot changed the title RFC: Shaped annotations that other type checkers can read Shaped annotations that other type checkers can read [rfc] (#5027) Sep 28, 2026
@github-actions github-actions Bot added size/xl and removed size/xl labels Sep 28, 2026
@codspeed

codspeed Bot commented Sep 28, 2026

Copy link
Copy Markdown

Merging this PR will improve performance by 54.03%

⚡ 1 improved benchmark
✅ 34 untouched benchmarks
⏩ 9 skipped benchmarks1

Performance Changes

Benchmark BASE HEAD Efficiency
⚡ invalidate_find 6.7 ms 4.3 ms +54.03%

Tip

Curious why performance improved? Comment @codspeedbot explain why performance improved on this PR, or directly use the CodSpeed MCP with your agent.


Comparing export-D122076108 (4f7a5e7) with main (39a5def)

Open in CodSpeed

Footnotes

  1. 9 benchmarks were skipped, so the baseline results were used instead. If they were deleted from the codebase, click here and archive them to remove them from the performance reports. ↩

@github-actions

Copy link
Copy Markdown

According to mypy_primer, this change doesn't affect type check results on a corpus of open source code. ✅

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant