From 63bb0fa98aea5bd4060bc240227f711928f73939 Mon Sep 17 00:00:00 2001 From: nstarman Date: Mon, 17 Aug 2026 15:46:42 -0400 Subject: [PATCH] fix: use typing_extensions backports so aliases resolve on Python 3.10/3.11 The predefined array aliases in bearshape.jax/numpy/torch/cupy were built with typing.TypeAliasType (3.12+) and typing.TypeVarTuple (3.11+), while the package declares requires-python = ">=3.10". Any checker resolving 3.10 or 3.11 reported "Variable not allowed in type expression" on every subscript of Shaped, F32, IntLike, and friends. Switch those aliases to the typing_extensions backports, and do the same for typing.Self in bearshape.cupy and typing.Never in bearshape._dimensions, which are 3.11+ and fail at the declared floor for the same reason. Add typing_extensions>=4.6 as a runtime dependency. Fixes #9 Co-Authored-By: Claude Opus 5 --- CHANGELOG.md | 11 +++ pyproject.toml | 5 +- src/bearshape/_dimensions.py | 4 +- src/bearshape/cupy.py | 108 ++++++++++++++-------------- src/bearshape/jax.py | 122 ++++++++++++++++---------------- src/bearshape/numpy.py | 133 ++++++++++++++++------------------- src/bearshape/torch.py | 122 ++++++++++++++++---------------- uv.lock | 6 +- 8 files changed, 263 insertions(+), 248 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 97d226e..7e038df 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,6 +13,17 @@ and this project follows - A GitHub Actions workflow for trusted publishing to PyPI, with automatic release-based publishing and manual `workflow_dispatch` support for a chosen ref. +- `typing_extensions>=4.6` as a runtime dependency, so the typing constructs + used by the backend aliases resolve on every supported Python version. + +### Fixed + +- Backend array aliases (`Shaped`, `F32`, `IntLike`, …) no longer break type + checkers resolving Python 3.10 or 3.11. The aliases used + `typing.TypeAliasType` (3.12+) and `typing.TypeVarTuple` (3.11+); they now + use the `typing_extensions` backports, along with `typing_extensions.Self` + in `bearshape.cupy` and `typing_extensions.Never` in + `bearshape._dimensions`. ### Changed diff --git a/pyproject.toml b/pyproject.toml index 860eb8d..ba1f45d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,7 +6,10 @@ readme = "README.md" license = "MIT" authors = [{ name = "acecchini", email = "ale.cecchini.valette@gmail.com" }] requires-python = ">=3.10" -dependencies = ["beartype>=0.20,<0.23"] # Tested with 0.20-0.22; see tox.toml +dependencies = [ + "beartype>=0.20,<0.23", # Tested with 0.20-0.22; see tox.toml + "typing_extensions>=4.6", +] keywords = [ "type-checking", "shape", diff --git a/src/bearshape/_dimensions.py b/src/bearshape/_dimensions.py index 1881815..d0768c2 100644 --- a/src/bearshape/_dimensions.py +++ b/src/bearshape/_dimensions.py @@ -36,6 +36,8 @@ def flatten(x: F32[N, C]) -> F32[N * C]: ... import typing as tp +from typing_extensions import Never + from ._shape import ( ANONYMOUS, ANONYMOUS_VARIADIC, @@ -312,7 +314,7 @@ def __pos__(self) -> _ValueExpr: def __neg__(self) -> _ValueExpr: return _ValueExpr(f"(-{self})") - def __invert__(self) -> tp.Never: + def __invert__(self) -> Never: """``~Value(...)`` is not supported — variadic requires a name.""" msg = "Value expressions cannot be variadic (~); use a named Dimension instead" raise TypeError(msg) diff --git a/src/bearshape/cupy.py b/src/bearshape/cupy.py index 5aad70b..b1806bb 100644 --- a/src/bearshape/cupy.py +++ b/src/bearshape/cupy.py @@ -19,6 +19,8 @@ def forward(x: F32[N, C, H, W]) -> F32[N, C, H, W]: ... import typing as tp +from typing_extensions import Self, TypeAliasType, TypeVarTuple + from ._imports import require_attr, require_module _CUPY_INSTALL_HINT = ( @@ -38,7 +40,7 @@ class CuPyArray(tp.Protocol): shape: tuple[int, ...] dtype: object - def __add__(self, other: object, /) -> tp.Self: ... + def __add__(self, other: object, /) -> Self: ... else: CuPyArray = tp.cast( @@ -232,36 +234,36 @@ def make_array_like_type( # --------------------------------------------------------------------------- if tp.TYPE_CHECKING: - _Dims = tp.TypeVarTuple("_Dims") + _Dims = TypeVarTuple("_Dims") - Bool = tp.TypeAliasType("Bool", CuPyArray, type_params=(_Dims,)) + Bool = TypeAliasType("Bool", CuPyArray, type_params=(_Dims,)) - I8 = tp.TypeAliasType("I8", CuPyArray, type_params=(_Dims,)) - I16 = tp.TypeAliasType("I16", CuPyArray, type_params=(_Dims,)) - I32 = tp.TypeAliasType("I32", CuPyArray, type_params=(_Dims,)) - I64 = tp.TypeAliasType("I64", CuPyArray, type_params=(_Dims,)) + I8 = TypeAliasType("I8", CuPyArray, type_params=(_Dims,)) + I16 = TypeAliasType("I16", CuPyArray, type_params=(_Dims,)) + I32 = TypeAliasType("I32", CuPyArray, type_params=(_Dims,)) + I64 = TypeAliasType("I64", CuPyArray, type_params=(_Dims,)) - U8 = tp.TypeAliasType("U8", CuPyArray, type_params=(_Dims,)) - U16 = tp.TypeAliasType("U16", CuPyArray, type_params=(_Dims,)) - U32 = tp.TypeAliasType("U32", CuPyArray, type_params=(_Dims,)) - U64 = tp.TypeAliasType("U64", CuPyArray, type_params=(_Dims,)) + U8 = TypeAliasType("U8", CuPyArray, type_params=(_Dims,)) + U16 = TypeAliasType("U16", CuPyArray, type_params=(_Dims,)) + U32 = TypeAliasType("U32", CuPyArray, type_params=(_Dims,)) + U64 = TypeAliasType("U64", CuPyArray, type_params=(_Dims,)) - F16 = tp.TypeAliasType("F16", CuPyArray, type_params=(_Dims,)) - F32 = tp.TypeAliasType("F32", CuPyArray, type_params=(_Dims,)) - F64 = tp.TypeAliasType("F64", CuPyArray, type_params=(_Dims,)) + F16 = TypeAliasType("F16", CuPyArray, type_params=(_Dims,)) + F32 = TypeAliasType("F32", CuPyArray, type_params=(_Dims,)) + F64 = TypeAliasType("F64", CuPyArray, type_params=(_Dims,)) - C64 = tp.TypeAliasType("C64", CuPyArray, type_params=(_Dims,)) - C128 = tp.TypeAliasType("C128", CuPyArray, type_params=(_Dims,)) + C64 = TypeAliasType("C64", CuPyArray, type_params=(_Dims,)) + C128 = TypeAliasType("C128", CuPyArray, type_params=(_Dims,)) - Int = tp.TypeAliasType("Int", CuPyArray, type_params=(_Dims,)) - UInt = tp.TypeAliasType("UInt", CuPyArray, type_params=(_Dims,)) - Integer = tp.TypeAliasType("Integer", CuPyArray, type_params=(_Dims,)) - Float = tp.TypeAliasType("Float", CuPyArray, type_params=(_Dims,)) - Real = tp.TypeAliasType("Real", CuPyArray, type_params=(_Dims,)) - Complex = tp.TypeAliasType("Complex", CuPyArray, type_params=(_Dims,)) - Inexact = tp.TypeAliasType("Inexact", CuPyArray, type_params=(_Dims,)) - Num = tp.TypeAliasType("Num", CuPyArray, type_params=(_Dims,)) - Shaped = tp.TypeAliasType("Shaped", CuPyArray, type_params=(_Dims,)) + Int = TypeAliasType("Int", CuPyArray, type_params=(_Dims,)) + UInt = TypeAliasType("UInt", CuPyArray, type_params=(_Dims,)) + Integer = TypeAliasType("Integer", CuPyArray, type_params=(_Dims,)) + Float = TypeAliasType("Float", CuPyArray, type_params=(_Dims,)) + Real = TypeAliasType("Real", CuPyArray, type_params=(_Dims,)) + Complex = TypeAliasType("Complex", CuPyArray, type_params=(_Dims,)) + Inexact = TypeAliasType("Inexact", CuPyArray, type_params=(_Dims,)) + Num = TypeAliasType("Num", CuPyArray, type_params=(_Dims,)) + Shaped = TypeAliasType("Shaped", CuPyArray, type_params=(_Dims,)) else: Bool = make_array_type(CuPyArray, BOOL) @@ -298,34 +300,34 @@ def make_array_like_type( # --------------------------------------------------------------------------- if tp.TYPE_CHECKING: - BoolLike = tp.TypeAliasType("BoolLike", CuPyArray, type_params=(_Dims,)) - - I8Like = tp.TypeAliasType("I8Like", CuPyArray, type_params=(_Dims,)) - I16Like = tp.TypeAliasType("I16Like", CuPyArray, type_params=(_Dims,)) - I32Like = tp.TypeAliasType("I32Like", CuPyArray, type_params=(_Dims,)) - I64Like = tp.TypeAliasType("I64Like", CuPyArray, type_params=(_Dims,)) - - U8Like = tp.TypeAliasType("U8Like", CuPyArray, type_params=(_Dims,)) - U16Like = tp.TypeAliasType("U16Like", CuPyArray, type_params=(_Dims,)) - U32Like = tp.TypeAliasType("U32Like", CuPyArray, type_params=(_Dims,)) - U64Like = tp.TypeAliasType("U64Like", CuPyArray, type_params=(_Dims,)) - - F16Like = tp.TypeAliasType("F16Like", CuPyArray, type_params=(_Dims,)) - F32Like = tp.TypeAliasType("F32Like", CuPyArray, type_params=(_Dims,)) - F64Like = tp.TypeAliasType("F64Like", CuPyArray, type_params=(_Dims,)) - - C64Like = tp.TypeAliasType("C64Like", CuPyArray, type_params=(_Dims,)) - C128Like = tp.TypeAliasType("C128Like", CuPyArray, type_params=(_Dims,)) - - IntLike = tp.TypeAliasType("IntLike", CuPyArray, type_params=(_Dims,)) - UIntLike = tp.TypeAliasType("UIntLike", CuPyArray, type_params=(_Dims,)) - IntegerLike = tp.TypeAliasType("IntegerLike", CuPyArray, type_params=(_Dims,)) - FloatLike = tp.TypeAliasType("FloatLike", CuPyArray, type_params=(_Dims,)) - RealLike = tp.TypeAliasType("RealLike", CuPyArray, type_params=(_Dims,)) - ComplexLike = tp.TypeAliasType("ComplexLike", CuPyArray, type_params=(_Dims,)) - InexactLike = tp.TypeAliasType("InexactLike", CuPyArray, type_params=(_Dims,)) - NumLike = tp.TypeAliasType("NumLike", CuPyArray, type_params=(_Dims,)) - ShapedLike = tp.TypeAliasType("ShapedLike", CuPyArray, type_params=(_Dims,)) + BoolLike = TypeAliasType("BoolLike", CuPyArray, type_params=(_Dims,)) + + I8Like = TypeAliasType("I8Like", CuPyArray, type_params=(_Dims,)) + I16Like = TypeAliasType("I16Like", CuPyArray, type_params=(_Dims,)) + I32Like = TypeAliasType("I32Like", CuPyArray, type_params=(_Dims,)) + I64Like = TypeAliasType("I64Like", CuPyArray, type_params=(_Dims,)) + + U8Like = TypeAliasType("U8Like", CuPyArray, type_params=(_Dims,)) + U16Like = TypeAliasType("U16Like", CuPyArray, type_params=(_Dims,)) + U32Like = TypeAliasType("U32Like", CuPyArray, type_params=(_Dims,)) + U64Like = TypeAliasType("U64Like", CuPyArray, type_params=(_Dims,)) + + F16Like = TypeAliasType("F16Like", CuPyArray, type_params=(_Dims,)) + F32Like = TypeAliasType("F32Like", CuPyArray, type_params=(_Dims,)) + F64Like = TypeAliasType("F64Like", CuPyArray, type_params=(_Dims,)) + + C64Like = TypeAliasType("C64Like", CuPyArray, type_params=(_Dims,)) + C128Like = TypeAliasType("C128Like", CuPyArray, type_params=(_Dims,)) + + IntLike = TypeAliasType("IntLike", CuPyArray, type_params=(_Dims,)) + UIntLike = TypeAliasType("UIntLike", CuPyArray, type_params=(_Dims,)) + IntegerLike = TypeAliasType("IntegerLike", CuPyArray, type_params=(_Dims,)) + FloatLike = TypeAliasType("FloatLike", CuPyArray, type_params=(_Dims,)) + RealLike = TypeAliasType("RealLike", CuPyArray, type_params=(_Dims,)) + ComplexLike = TypeAliasType("ComplexLike", CuPyArray, type_params=(_Dims,)) + InexactLike = TypeAliasType("InexactLike", CuPyArray, type_params=(_Dims,)) + NumLike = TypeAliasType("NumLike", CuPyArray, type_params=(_Dims,)) + ShapedLike = TypeAliasType("ShapedLike", CuPyArray, type_params=(_Dims,)) else: BoolLike = make_array_like_type(BOOL, name="BoolLike") diff --git a/src/bearshape/jax.py b/src/bearshape/jax.py index b4eb711..8ac0cb1 100644 --- a/src/bearshape/jax.py +++ b/src/bearshape/jax.py @@ -19,6 +19,8 @@ def forward(x: F32[N, C, H, W]) -> BF16[N, C, H, W]: ... import typing as tp +from typing_extensions import TypeAliasType, TypeVarTuple + from ._imports import require_attr, require_module _JAX_INSTALL_HINT = ( @@ -232,37 +234,37 @@ def make_array_like_type( # --------------------------------------------------------------------------- if tp.TYPE_CHECKING: - _Dims = tp.TypeVarTuple("_Dims") - - Bool = tp.TypeAliasType("Bool", JaxArray, type_params=(_Dims,)) - - I8 = tp.TypeAliasType("I8", JaxArray, type_params=(_Dims,)) - I16 = tp.TypeAliasType("I16", JaxArray, type_params=(_Dims,)) - I32 = tp.TypeAliasType("I32", JaxArray, type_params=(_Dims,)) - I64 = tp.TypeAliasType("I64", JaxArray, type_params=(_Dims,)) - - U8 = tp.TypeAliasType("U8", JaxArray, type_params=(_Dims,)) - U16 = tp.TypeAliasType("U16", JaxArray, type_params=(_Dims,)) - U32 = tp.TypeAliasType("U32", JaxArray, type_params=(_Dims,)) - U64 = tp.TypeAliasType("U64", JaxArray, type_params=(_Dims,)) - - F16 = tp.TypeAliasType("F16", JaxArray, type_params=(_Dims,)) - F32 = tp.TypeAliasType("F32", JaxArray, type_params=(_Dims,)) - F64 = tp.TypeAliasType("F64", JaxArray, type_params=(_Dims,)) - BF16 = tp.TypeAliasType("BF16", JaxArray, type_params=(_Dims,)) - - C64 = tp.TypeAliasType("C64", JaxArray, type_params=(_Dims,)) - C128 = tp.TypeAliasType("C128", JaxArray, type_params=(_Dims,)) - - Int = tp.TypeAliasType("Int", JaxArray, type_params=(_Dims,)) - UInt = tp.TypeAliasType("UInt", JaxArray, type_params=(_Dims,)) - Integer = tp.TypeAliasType("Integer", JaxArray, type_params=(_Dims,)) - Float = tp.TypeAliasType("Float", JaxArray, type_params=(_Dims,)) - Real = tp.TypeAliasType("Real", JaxArray, type_params=(_Dims,)) - Complex = tp.TypeAliasType("Complex", JaxArray, type_params=(_Dims,)) - Inexact = tp.TypeAliasType("Inexact", JaxArray, type_params=(_Dims,)) - Num = tp.TypeAliasType("Num", JaxArray, type_params=(_Dims,)) - Shaped = tp.TypeAliasType("Shaped", JaxArray, type_params=(_Dims,)) + _Dims = TypeVarTuple("_Dims") + + Bool = TypeAliasType("Bool", JaxArray, type_params=(_Dims,)) + + I8 = TypeAliasType("I8", JaxArray, type_params=(_Dims,)) + I16 = TypeAliasType("I16", JaxArray, type_params=(_Dims,)) + I32 = TypeAliasType("I32", JaxArray, type_params=(_Dims,)) + I64 = TypeAliasType("I64", JaxArray, type_params=(_Dims,)) + + U8 = TypeAliasType("U8", JaxArray, type_params=(_Dims,)) + U16 = TypeAliasType("U16", JaxArray, type_params=(_Dims,)) + U32 = TypeAliasType("U32", JaxArray, type_params=(_Dims,)) + U64 = TypeAliasType("U64", JaxArray, type_params=(_Dims,)) + + F16 = TypeAliasType("F16", JaxArray, type_params=(_Dims,)) + F32 = TypeAliasType("F32", JaxArray, type_params=(_Dims,)) + F64 = TypeAliasType("F64", JaxArray, type_params=(_Dims,)) + BF16 = TypeAliasType("BF16", JaxArray, type_params=(_Dims,)) + + C64 = TypeAliasType("C64", JaxArray, type_params=(_Dims,)) + C128 = TypeAliasType("C128", JaxArray, type_params=(_Dims,)) + + Int = TypeAliasType("Int", JaxArray, type_params=(_Dims,)) + UInt = TypeAliasType("UInt", JaxArray, type_params=(_Dims,)) + Integer = TypeAliasType("Integer", JaxArray, type_params=(_Dims,)) + Float = TypeAliasType("Float", JaxArray, type_params=(_Dims,)) + Real = TypeAliasType("Real", JaxArray, type_params=(_Dims,)) + Complex = TypeAliasType("Complex", JaxArray, type_params=(_Dims,)) + Inexact = TypeAliasType("Inexact", JaxArray, type_params=(_Dims,)) + Num = TypeAliasType("Num", JaxArray, type_params=(_Dims,)) + Shaped = TypeAliasType("Shaped", JaxArray, type_params=(_Dims,)) else: Bool = make_array_type(JaxArray, BOOL) @@ -300,35 +302,35 @@ def make_array_like_type( # --------------------------------------------------------------------------- if tp.TYPE_CHECKING: - BF16Like = tp.TypeAliasType("BF16Like", JaxArray, type_params=(_Dims,)) - BoolLike = tp.TypeAliasType("BoolLike", JaxArray, type_params=(_Dims,)) - - I8Like = tp.TypeAliasType("I8Like", JaxArray, type_params=(_Dims,)) - I16Like = tp.TypeAliasType("I16Like", JaxArray, type_params=(_Dims,)) - I32Like = tp.TypeAliasType("I32Like", JaxArray, type_params=(_Dims,)) - I64Like = tp.TypeAliasType("I64Like", JaxArray, type_params=(_Dims,)) - - U8Like = tp.TypeAliasType("U8Like", JaxArray, type_params=(_Dims,)) - U16Like = tp.TypeAliasType("U16Like", JaxArray, type_params=(_Dims,)) - U32Like = tp.TypeAliasType("U32Like", JaxArray, type_params=(_Dims,)) - U64Like = tp.TypeAliasType("U64Like", JaxArray, type_params=(_Dims,)) - - F16Like = tp.TypeAliasType("F16Like", JaxArray, type_params=(_Dims,)) - F32Like = tp.TypeAliasType("F32Like", JaxArray, type_params=(_Dims,)) - F64Like = tp.TypeAliasType("F64Like", JaxArray, type_params=(_Dims,)) - - C64Like = tp.TypeAliasType("C64Like", JaxArray, type_params=(_Dims,)) - C128Like = tp.TypeAliasType("C128Like", JaxArray, type_params=(_Dims,)) - - IntLike = tp.TypeAliasType("IntLike", JaxArray, type_params=(_Dims,)) - UIntLike = tp.TypeAliasType("UIntLike", JaxArray, type_params=(_Dims,)) - IntegerLike = tp.TypeAliasType("IntegerLike", JaxArray, type_params=(_Dims,)) - FloatLike = tp.TypeAliasType("FloatLike", JaxArray, type_params=(_Dims,)) - RealLike = tp.TypeAliasType("RealLike", JaxArray, type_params=(_Dims,)) - ComplexLike = tp.TypeAliasType("ComplexLike", JaxArray, type_params=(_Dims,)) - InexactLike = tp.TypeAliasType("InexactLike", JaxArray, type_params=(_Dims,)) - NumLike = tp.TypeAliasType("NumLike", JaxArray, type_params=(_Dims,)) - ShapedLike = tp.TypeAliasType("ShapedLike", JaxArray, type_params=(_Dims,)) + BF16Like = TypeAliasType("BF16Like", JaxArray, type_params=(_Dims,)) + BoolLike = TypeAliasType("BoolLike", JaxArray, type_params=(_Dims,)) + + I8Like = TypeAliasType("I8Like", JaxArray, type_params=(_Dims,)) + I16Like = TypeAliasType("I16Like", JaxArray, type_params=(_Dims,)) + I32Like = TypeAliasType("I32Like", JaxArray, type_params=(_Dims,)) + I64Like = TypeAliasType("I64Like", JaxArray, type_params=(_Dims,)) + + U8Like = TypeAliasType("U8Like", JaxArray, type_params=(_Dims,)) + U16Like = TypeAliasType("U16Like", JaxArray, type_params=(_Dims,)) + U32Like = TypeAliasType("U32Like", JaxArray, type_params=(_Dims,)) + U64Like = TypeAliasType("U64Like", JaxArray, type_params=(_Dims,)) + + F16Like = TypeAliasType("F16Like", JaxArray, type_params=(_Dims,)) + F32Like = TypeAliasType("F32Like", JaxArray, type_params=(_Dims,)) + F64Like = TypeAliasType("F64Like", JaxArray, type_params=(_Dims,)) + + C64Like = TypeAliasType("C64Like", JaxArray, type_params=(_Dims,)) + C128Like = TypeAliasType("C128Like", JaxArray, type_params=(_Dims,)) + + IntLike = TypeAliasType("IntLike", JaxArray, type_params=(_Dims,)) + UIntLike = TypeAliasType("UIntLike", JaxArray, type_params=(_Dims,)) + IntegerLike = TypeAliasType("IntegerLike", JaxArray, type_params=(_Dims,)) + FloatLike = TypeAliasType("FloatLike", JaxArray, type_params=(_Dims,)) + RealLike = TypeAliasType("RealLike", JaxArray, type_params=(_Dims,)) + ComplexLike = TypeAliasType("ComplexLike", JaxArray, type_params=(_Dims,)) + InexactLike = TypeAliasType("InexactLike", JaxArray, type_params=(_Dims,)) + NumLike = TypeAliasType("NumLike", JaxArray, type_params=(_Dims,)) + ShapedLike = TypeAliasType("ShapedLike", JaxArray, type_params=(_Dims,)) else: BF16Like = make_array_like_type(BFLOAT16, name="BF16Like") diff --git a/src/bearshape/numpy.py b/src/bearshape/numpy.py index c79eecc..86ca13b 100644 --- a/src/bearshape/numpy.py +++ b/src/bearshape/numpy.py @@ -58,6 +58,7 @@ def pixel(value: U8ScalarLike) -> int: ... # [0, 255] import numpy as np from beartype.vale import Is from numpy._typing import _NestedSequence, _SupportsArray +from typing_extensions import TypeAliasType, TypeVarTuple __all__ = [ # Array types — base @@ -368,127 +369,115 @@ def f(points: Point[N]) -> Point[N]: ... if tp.TYPE_CHECKING: from numpy.typing import NDArray - _Dims = tp.TypeVarTuple("_Dims") + _Dims = TypeVarTuple("_Dims") # --- Base types --- - Bool = tp.TypeAliasType("Bool", NDArray[np.bool_], type_params=(_Dims,)) + Bool = TypeAliasType("Bool", NDArray[np.bool_], type_params=(_Dims,)) - I8 = tp.TypeAliasType("I8", NDArray[np.int8], type_params=(_Dims,)) - I16 = tp.TypeAliasType("I16", NDArray[np.int16], type_params=(_Dims,)) - I32 = tp.TypeAliasType("I32", NDArray[np.int32], type_params=(_Dims,)) - I64 = tp.TypeAliasType("I64", NDArray[np.int64], type_params=(_Dims,)) + I8 = TypeAliasType("I8", NDArray[np.int8], type_params=(_Dims,)) + I16 = TypeAliasType("I16", NDArray[np.int16], type_params=(_Dims,)) + I32 = TypeAliasType("I32", NDArray[np.int32], type_params=(_Dims,)) + I64 = TypeAliasType("I64", NDArray[np.int64], type_params=(_Dims,)) - U8 = tp.TypeAliasType("U8", NDArray[np.uint8], type_params=(_Dims,)) - U16 = tp.TypeAliasType("U16", NDArray[np.uint16], type_params=(_Dims,)) - U32 = tp.TypeAliasType("U32", NDArray[np.uint32], type_params=(_Dims,)) - U64 = tp.TypeAliasType("U64", NDArray[np.uint64], type_params=(_Dims,)) + U8 = TypeAliasType("U8", NDArray[np.uint8], type_params=(_Dims,)) + U16 = TypeAliasType("U16", NDArray[np.uint16], type_params=(_Dims,)) + U32 = TypeAliasType("U32", NDArray[np.uint32], type_params=(_Dims,)) + U64 = TypeAliasType("U64", NDArray[np.uint64], type_params=(_Dims,)) - F16 = tp.TypeAliasType("F16", NDArray[np.float16], type_params=(_Dims,)) - F32 = tp.TypeAliasType("F32", NDArray[np.float32], type_params=(_Dims,)) - F64 = tp.TypeAliasType("F64", NDArray[np.float64], type_params=(_Dims,)) - F128 = tp.TypeAliasType("F128", NDArray[np.longdouble], type_params=(_Dims,)) + F16 = TypeAliasType("F16", NDArray[np.float16], type_params=(_Dims,)) + F32 = TypeAliasType("F32", NDArray[np.float32], type_params=(_Dims,)) + F64 = TypeAliasType("F64", NDArray[np.float64], type_params=(_Dims,)) + F128 = TypeAliasType("F128", NDArray[np.longdouble], type_params=(_Dims,)) - C64 = tp.TypeAliasType("C64", NDArray[np.complex64], type_params=(_Dims,)) - C128 = tp.TypeAliasType("C128", NDArray[np.complex128], type_params=(_Dims,)) - C256 = tp.TypeAliasType("C256", NDArray[np.clongdouble], type_params=(_Dims,)) + C64 = TypeAliasType("C64", NDArray[np.complex64], type_params=(_Dims,)) + C128 = TypeAliasType("C128", NDArray[np.complex128], type_params=(_Dims,)) + C256 = TypeAliasType("C256", NDArray[np.clongdouble], type_params=(_Dims,)) - Int = tp.TypeAliasType("Int", NDArray[np.signedinteger[tp.Any]], type_params=(_Dims,)) - UInt = tp.TypeAliasType( + Int = TypeAliasType("Int", NDArray[np.signedinteger[tp.Any]], type_params=(_Dims,)) + UInt = TypeAliasType( "UInt", NDArray[np.unsignedinteger[tp.Any]], type_params=(_Dims,) ) - Integer = tp.TypeAliasType( - "Integer", NDArray[np.integer[tp.Any]], type_params=(_Dims,) - ) - Float = tp.TypeAliasType("Float", NDArray[np.floating[tp.Any]], type_params=(_Dims,)) - Real = tp.TypeAliasType( + Integer = TypeAliasType("Integer", NDArray[np.integer[tp.Any]], type_params=(_Dims,)) + Float = TypeAliasType("Float", NDArray[np.floating[tp.Any]], type_params=(_Dims,)) + Real = TypeAliasType( "Real", NDArray[np.integer[tp.Any] | np.floating[tp.Any]], type_params=(_Dims,) ) - Complex = tp.TypeAliasType( + Complex = TypeAliasType( "Complex", NDArray[np.complexfloating[tp.Any, tp.Any]], type_params=(_Dims,) ) - Inexact = tp.TypeAliasType( - "Inexact", NDArray[np.inexact[tp.Any]], type_params=(_Dims,) - ) - Num = tp.TypeAliasType("Num", NDArray[np.number[tp.Any]], type_params=(_Dims,)) - Shaped = tp.TypeAliasType( + Inexact = TypeAliasType("Inexact", NDArray[np.inexact[tp.Any]], type_params=(_Dims,)) + Num = TypeAliasType("Num", NDArray[np.number[tp.Any]], type_params=(_Dims,)) + Shaped = TypeAliasType( "Shaped", NDArray[np.bool_ | np.number[tp.Any]], type_params=(_Dims,) ) # --- New dtypes --- - V = tp.TypeAliasType("V", NDArray[np.void], type_params=(_Dims,)) - Str = tp.TypeAliasType("Str", NDArray[np.str_], type_params=(_Dims,)) - Bytes = tp.TypeAliasType("Bytes", NDArray[np.bytes_], type_params=(_Dims,)) - Obj = tp.TypeAliasType("Obj", NDArray[np.object_], type_params=(_Dims,)) - DT64 = tp.TypeAliasType("DT64", NDArray[np.datetime64], type_params=(_Dims,)) - TD64 = tp.TypeAliasType("TD64", NDArray[np.timedelta64], type_params=(_Dims,)) + V = TypeAliasType("V", NDArray[np.void], type_params=(_Dims,)) + Str = TypeAliasType("Str", NDArray[np.str_], type_params=(_Dims,)) + Bytes = TypeAliasType("Bytes", NDArray[np.bytes_], type_params=(_Dims,)) + Obj = TypeAliasType("Obj", NDArray[np.object_], type_params=(_Dims,)) + DT64 = TypeAliasType("DT64", NDArray[np.datetime64], type_params=(_Dims,)) + TD64 = TypeAliasType("TD64", NDArray[np.timedelta64], type_params=(_Dims,)) # --- Like types (static: ArrayLike template with bare scalar types) --- - BoolLike = tp.TypeAliasType( - "BoolLike", ArrayLike[bool, np.bool_], type_params=(_Dims,) - ) - - I8Like = tp.TypeAliasType("I8Like", ArrayLike[int, np.int8], type_params=(_Dims,)) - I16Like = tp.TypeAliasType("I16Like", ArrayLike[int, np.int16], type_params=(_Dims,)) - I32Like = tp.TypeAliasType("I32Like", ArrayLike[int, np.int32], type_params=(_Dims,)) - I64Like = tp.TypeAliasType("I64Like", ArrayLike[int, np.int64], type_params=(_Dims,)) - - U8Like = tp.TypeAliasType("U8Like", ArrayLike[int, np.uint8], type_params=(_Dims,)) - U16Like = tp.TypeAliasType("U16Like", ArrayLike[int, np.uint16], type_params=(_Dims,)) - U32Like = tp.TypeAliasType("U32Like", ArrayLike[int, np.uint32], type_params=(_Dims,)) - U64Like = tp.TypeAliasType("U64Like", ArrayLike[int, np.uint64], type_params=(_Dims,)) - - F16Like = tp.TypeAliasType( - "F16Like", ArrayLike[float, np.float16], type_params=(_Dims,) - ) - F32Like = tp.TypeAliasType( - "F32Like", ArrayLike[float, np.float32], type_params=(_Dims,) - ) - F64Like = tp.TypeAliasType( - "F64Like", ArrayLike[float, np.float64], type_params=(_Dims,) - ) - F128Like = tp.TypeAliasType( + BoolLike = TypeAliasType("BoolLike", ArrayLike[bool, np.bool_], type_params=(_Dims,)) + + I8Like = TypeAliasType("I8Like", ArrayLike[int, np.int8], type_params=(_Dims,)) + I16Like = TypeAliasType("I16Like", ArrayLike[int, np.int16], type_params=(_Dims,)) + I32Like = TypeAliasType("I32Like", ArrayLike[int, np.int32], type_params=(_Dims,)) + I64Like = TypeAliasType("I64Like", ArrayLike[int, np.int64], type_params=(_Dims,)) + + U8Like = TypeAliasType("U8Like", ArrayLike[int, np.uint8], type_params=(_Dims,)) + U16Like = TypeAliasType("U16Like", ArrayLike[int, np.uint16], type_params=(_Dims,)) + U32Like = TypeAliasType("U32Like", ArrayLike[int, np.uint32], type_params=(_Dims,)) + U64Like = TypeAliasType("U64Like", ArrayLike[int, np.uint64], type_params=(_Dims,)) + + F16Like = TypeAliasType("F16Like", ArrayLike[float, np.float16], type_params=(_Dims,)) + F32Like = TypeAliasType("F32Like", ArrayLike[float, np.float32], type_params=(_Dims,)) + F64Like = TypeAliasType("F64Like", ArrayLike[float, np.float64], type_params=(_Dims,)) + F128Like = TypeAliasType( "F128Like", ArrayLike[float, np.longdouble], type_params=(_Dims,) ) - C64Like = tp.TypeAliasType( + C64Like = TypeAliasType( "C64Like", ArrayLike[complex, np.complex64], type_params=(_Dims,) ) - C128Like = tp.TypeAliasType( + C128Like = TypeAliasType( "C128Like", ArrayLike[complex, np.complex128], type_params=(_Dims,) ) - C256Like = tp.TypeAliasType( + C256Like = TypeAliasType( "C256Like", ArrayLike[complex, np.clongdouble], type_params=(_Dims,) ) - IntLike = tp.TypeAliasType( + IntLike = TypeAliasType( "IntLike", ArrayLike[int, np.signedinteger[tp.Any]], type_params=(_Dims,) ) - UIntLike = tp.TypeAliasType( + UIntLike = TypeAliasType( "UIntLike", ArrayLike[int, np.unsignedinteger[tp.Any]], type_params=(_Dims,) ) - IntegerLike = tp.TypeAliasType( + IntegerLike = TypeAliasType( "IntegerLike", ArrayLike[int, np.integer[tp.Any]], type_params=(_Dims,) ) - FloatLike = tp.TypeAliasType( + FloatLike = TypeAliasType( "FloatLike", ArrayLike[float, np.floating[tp.Any]], type_params=(_Dims,) ) - RealLike = tp.TypeAliasType( + RealLike = TypeAliasType( "RealLike", ArrayLike[int | float, np.integer[tp.Any] | np.floating[tp.Any]], type_params=(_Dims,), ) - ComplexLike = tp.TypeAliasType( + ComplexLike = TypeAliasType( "ComplexLike", ArrayLike[complex, np.complexfloating[tp.Any, tp.Any]], type_params=(_Dims,), ) - InexactLike = tp.TypeAliasType( + InexactLike = TypeAliasType( "InexactLike", ArrayLike[float | complex, np.inexact[tp.Any]], type_params=(_Dims,) ) - NumLike = tp.TypeAliasType( + NumLike = TypeAliasType( "NumLike", ArrayLike[int | float | complex, np.number[tp.Any]], type_params=(_Dims,) ) - ShapedLike = tp.TypeAliasType( + ShapedLike = TypeAliasType( "ShapedLike", ArrayLike[bool | int | float | complex, np.bool_ | np.number[tp.Any]], type_params=(_Dims,), diff --git a/src/bearshape/torch.py b/src/bearshape/torch.py index 1309c65..40ca5ab 100644 --- a/src/bearshape/torch.py +++ b/src/bearshape/torch.py @@ -19,6 +19,8 @@ def forward(x: F32[N, C, H, W]) -> F32[N, C, H, W]: ... import typing as tp +from typing_extensions import TypeAliasType, TypeVarTuple + from ._imports import require_attr, require_module _TORCH_INSTALL_HINT = ( @@ -227,37 +229,37 @@ def make_array_like_type( # --------------------------------------------------------------------------- if tp.TYPE_CHECKING: - _Dims = tp.TypeVarTuple("_Dims") - - Bool = tp.TypeAliasType("Bool", Tensor, type_params=(_Dims,)) - - I8 = tp.TypeAliasType("I8", Tensor, type_params=(_Dims,)) - I16 = tp.TypeAliasType("I16", Tensor, type_params=(_Dims,)) - I32 = tp.TypeAliasType("I32", Tensor, type_params=(_Dims,)) - I64 = tp.TypeAliasType("I64", Tensor, type_params=(_Dims,)) - - U8 = tp.TypeAliasType("U8", Tensor, type_params=(_Dims,)) - U16 = tp.TypeAliasType("U16", Tensor, type_params=(_Dims,)) - U32 = tp.TypeAliasType("U32", Tensor, type_params=(_Dims,)) - U64 = tp.TypeAliasType("U64", Tensor, type_params=(_Dims,)) - - F16 = tp.TypeAliasType("F16", Tensor, type_params=(_Dims,)) - F32 = tp.TypeAliasType("F32", Tensor, type_params=(_Dims,)) - F64 = tp.TypeAliasType("F64", Tensor, type_params=(_Dims,)) - BF16 = tp.TypeAliasType("BF16", Tensor, type_params=(_Dims,)) - - C64 = tp.TypeAliasType("C64", Tensor, type_params=(_Dims,)) - C128 = tp.TypeAliasType("C128", Tensor, type_params=(_Dims,)) - - Int = tp.TypeAliasType("Int", Tensor, type_params=(_Dims,)) - UInt = tp.TypeAliasType("UInt", Tensor, type_params=(_Dims,)) - Integer = tp.TypeAliasType("Integer", Tensor, type_params=(_Dims,)) - Float = tp.TypeAliasType("Float", Tensor, type_params=(_Dims,)) - Real = tp.TypeAliasType("Real", Tensor, type_params=(_Dims,)) - Complex = tp.TypeAliasType("Complex", Tensor, type_params=(_Dims,)) - Inexact = tp.TypeAliasType("Inexact", Tensor, type_params=(_Dims,)) - Num = tp.TypeAliasType("Num", Tensor, type_params=(_Dims,)) - Shaped = tp.TypeAliasType("Shaped", Tensor, type_params=(_Dims,)) + _Dims = TypeVarTuple("_Dims") + + Bool = TypeAliasType("Bool", Tensor, type_params=(_Dims,)) + + I8 = TypeAliasType("I8", Tensor, type_params=(_Dims,)) + I16 = TypeAliasType("I16", Tensor, type_params=(_Dims,)) + I32 = TypeAliasType("I32", Tensor, type_params=(_Dims,)) + I64 = TypeAliasType("I64", Tensor, type_params=(_Dims,)) + + U8 = TypeAliasType("U8", Tensor, type_params=(_Dims,)) + U16 = TypeAliasType("U16", Tensor, type_params=(_Dims,)) + U32 = TypeAliasType("U32", Tensor, type_params=(_Dims,)) + U64 = TypeAliasType("U64", Tensor, type_params=(_Dims,)) + + F16 = TypeAliasType("F16", Tensor, type_params=(_Dims,)) + F32 = TypeAliasType("F32", Tensor, type_params=(_Dims,)) + F64 = TypeAliasType("F64", Tensor, type_params=(_Dims,)) + BF16 = TypeAliasType("BF16", Tensor, type_params=(_Dims,)) + + C64 = TypeAliasType("C64", Tensor, type_params=(_Dims,)) + C128 = TypeAliasType("C128", Tensor, type_params=(_Dims,)) + + Int = TypeAliasType("Int", Tensor, type_params=(_Dims,)) + UInt = TypeAliasType("UInt", Tensor, type_params=(_Dims,)) + Integer = TypeAliasType("Integer", Tensor, type_params=(_Dims,)) + Float = TypeAliasType("Float", Tensor, type_params=(_Dims,)) + Real = TypeAliasType("Real", Tensor, type_params=(_Dims,)) + Complex = TypeAliasType("Complex", Tensor, type_params=(_Dims,)) + Inexact = TypeAliasType("Inexact", Tensor, type_params=(_Dims,)) + Num = TypeAliasType("Num", Tensor, type_params=(_Dims,)) + Shaped = TypeAliasType("Shaped", Tensor, type_params=(_Dims,)) else: Bool = make_array_type(Tensor, BOOL) @@ -295,35 +297,35 @@ def make_array_like_type( # --------------------------------------------------------------------------- if tp.TYPE_CHECKING: - BF16Like = tp.TypeAliasType("BF16Like", Tensor, type_params=(_Dims,)) - BoolLike = tp.TypeAliasType("BoolLike", Tensor, type_params=(_Dims,)) - - I8Like = tp.TypeAliasType("I8Like", Tensor, type_params=(_Dims,)) - I16Like = tp.TypeAliasType("I16Like", Tensor, type_params=(_Dims,)) - I32Like = tp.TypeAliasType("I32Like", Tensor, type_params=(_Dims,)) - I64Like = tp.TypeAliasType("I64Like", Tensor, type_params=(_Dims,)) - - U8Like = tp.TypeAliasType("U8Like", Tensor, type_params=(_Dims,)) - U16Like = tp.TypeAliasType("U16Like", Tensor, type_params=(_Dims,)) - U32Like = tp.TypeAliasType("U32Like", Tensor, type_params=(_Dims,)) - U64Like = tp.TypeAliasType("U64Like", Tensor, type_params=(_Dims,)) - - F16Like = tp.TypeAliasType("F16Like", Tensor, type_params=(_Dims,)) - F32Like = tp.TypeAliasType("F32Like", Tensor, type_params=(_Dims,)) - F64Like = tp.TypeAliasType("F64Like", Tensor, type_params=(_Dims,)) - - C64Like = tp.TypeAliasType("C64Like", Tensor, type_params=(_Dims,)) - C128Like = tp.TypeAliasType("C128Like", Tensor, type_params=(_Dims,)) - - IntLike = tp.TypeAliasType("IntLike", Tensor, type_params=(_Dims,)) - UIntLike = tp.TypeAliasType("UIntLike", Tensor, type_params=(_Dims,)) - IntegerLike = tp.TypeAliasType("IntegerLike", Tensor, type_params=(_Dims,)) - FloatLike = tp.TypeAliasType("FloatLike", Tensor, type_params=(_Dims,)) - RealLike = tp.TypeAliasType("RealLike", Tensor, type_params=(_Dims,)) - ComplexLike = tp.TypeAliasType("ComplexLike", Tensor, type_params=(_Dims,)) - InexactLike = tp.TypeAliasType("InexactLike", Tensor, type_params=(_Dims,)) - NumLike = tp.TypeAliasType("NumLike", Tensor, type_params=(_Dims,)) - ShapedLike = tp.TypeAliasType("ShapedLike", Tensor, type_params=(_Dims,)) + BF16Like = TypeAliasType("BF16Like", Tensor, type_params=(_Dims,)) + BoolLike = TypeAliasType("BoolLike", Tensor, type_params=(_Dims,)) + + I8Like = TypeAliasType("I8Like", Tensor, type_params=(_Dims,)) + I16Like = TypeAliasType("I16Like", Tensor, type_params=(_Dims,)) + I32Like = TypeAliasType("I32Like", Tensor, type_params=(_Dims,)) + I64Like = TypeAliasType("I64Like", Tensor, type_params=(_Dims,)) + + U8Like = TypeAliasType("U8Like", Tensor, type_params=(_Dims,)) + U16Like = TypeAliasType("U16Like", Tensor, type_params=(_Dims,)) + U32Like = TypeAliasType("U32Like", Tensor, type_params=(_Dims,)) + U64Like = TypeAliasType("U64Like", Tensor, type_params=(_Dims,)) + + F16Like = TypeAliasType("F16Like", Tensor, type_params=(_Dims,)) + F32Like = TypeAliasType("F32Like", Tensor, type_params=(_Dims,)) + F64Like = TypeAliasType("F64Like", Tensor, type_params=(_Dims,)) + + C64Like = TypeAliasType("C64Like", Tensor, type_params=(_Dims,)) + C128Like = TypeAliasType("C128Like", Tensor, type_params=(_Dims,)) + + IntLike = TypeAliasType("IntLike", Tensor, type_params=(_Dims,)) + UIntLike = TypeAliasType("UIntLike", Tensor, type_params=(_Dims,)) + IntegerLike = TypeAliasType("IntegerLike", Tensor, type_params=(_Dims,)) + FloatLike = TypeAliasType("FloatLike", Tensor, type_params=(_Dims,)) + RealLike = TypeAliasType("RealLike", Tensor, type_params=(_Dims,)) + ComplexLike = TypeAliasType("ComplexLike", Tensor, type_params=(_Dims,)) + InexactLike = TypeAliasType("InexactLike", Tensor, type_params=(_Dims,)) + NumLike = TypeAliasType("NumLike", Tensor, type_params=(_Dims,)) + ShapedLike = TypeAliasType("ShapedLike", Tensor, type_params=(_Dims,)) else: BF16Like = make_array_like_type(BFLOAT16, name="BF16Like") diff --git a/uv.lock b/uv.lock index be754b2..12ce4a8 100644 --- a/uv.lock +++ b/uv.lock @@ -36,6 +36,7 @@ version = "0.0.1" source = { editable = "." } dependencies = [ { name = "beartype" }, + { name = "typing-extensions" }, ] [package.dev-dependencies] @@ -72,7 +73,10 @@ test = [ ] [package.metadata] -requires-dist = [{ name = "beartype", specifier = ">=0.20,<0.23" }] +requires-dist = [ + { name = "beartype", specifier = ">=0.20,<0.23" }, + { name = "typing-extensions", specifier = ">=4.6" }, +] [package.metadata.requires-dev] dev = [