Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
5 changes: 4 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,10 @@ readme = "README.md"
license = "MIT"
authors = [{ name = "acecchini", email = "[email protected]" }]
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",
Expand Down
4 changes: 3 additions & 1 deletion src/bearshape/_dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
108 changes: 55 additions & 53 deletions src/bearshape/cupy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = (
Expand All @@ -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(
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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")
Expand Down
122 changes: 62 additions & 60 deletions src/bearshape/jax.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = (
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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")
Expand Down
Loading