Changes from feedback:

- Float32 introduced to replace f32 (etc.)
- Old f32 aliases were previously around for backward-compatibility, but
  the code is pretty hideous, and all stakeholders are now on board
  through RFC #13.
- Overall feedback was that Int{Sign,Unsign} was too long and not enough
  like numpy. Changed to Int, UInt, and Integer.
This commit is contained in:
Patrick Kidger
2022-09-07 11:09:48 -07:00
parent 140be9ecec
commit a53fe6af57
8 changed files with 138 additions and 136 deletions
+2 -2
View File
@@ -20,12 +20,12 @@
import jax.numpy as jnp
import pytest
from jaxtyping import f32
from jaxtyping import Float32
from .helpers import ParamError
def g(x: f32[jnp.ndarray, " b"]):
def g(x: Float32[jnp.ndarray, " b"]):
pass
+2 -2
View File
@@ -20,12 +20,12 @@
import jax.numpy as jnp
import pytest
from jaxtyping import f32
from jaxtyping import Float32
from .helpers import ParamError
def g(x: f32[jnp.ndarray, " b"]):
def g(x: Float32[jnp.ndarray, " b"]):
pass
@@ -20,12 +20,12 @@
import jax.numpy as jnp
import pytest
from jaxtyping import f32
from jaxtyping import Float32
from ..helpers import ParamError
def g(x: f32[jnp.ndarray, " b"]):
def g(x: Float32[jnp.ndarray, " b"]):
pass
+2 -2
View File
@@ -20,12 +20,12 @@
import jax.numpy as jnp
import pytest
from jaxtyping import f32
from jaxtyping import Float32
from .helpers import ParamError
def g(x: f32[jnp.ndarray, " b"]):
def g(x: Float32[jnp.ndarray, " b"]):
pass
+57 -23
View File
@@ -21,7 +21,7 @@ import jax.numpy as jnp
import jax.random as jr
import pytest
from jaxtyping import Array, f32, Float, jaxtyped, Shaped
from jaxtyping import Array, Float, Float32, jaxtyped, Shaped
from .helpers import ParamError, ReturnError
@@ -35,6 +35,34 @@ def test_basic(typecheck):
g(jnp.array(1.0))
def test_dtypes():
from jaxtyping import ( # noqa: F401
Array,
BFloat16,
Bool,
Complex,
Complex64,
Complex128,
Float,
Float16,
Float32,
Float64,
Inexact,
Int,
Int8,
Int16,
Int32,
Int64,
Num,
Shaped,
UInt,
UInt8,
UInt16,
UInt32,
UInt64,
)
def test_return(typecheck, getkey):
@jaxtyped
@typecheck
@@ -92,12 +120,12 @@ def test_any_dtype(typecheck, getkey):
def test_nested_jaxtyped(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "b c"], transpose: bool) -> f32[Array, "c b"]:
def g(x: Float32[Array, "b c"], transpose: bool) -> Float32[Array, "c b"]:
return h(x, transpose)
@jaxtyped
@typecheck
def h(x: f32[Array, "c b"], transpose: bool) -> f32[Array, "b c"]:
def h(x: Float32[Array, "c b"], transpose: bool) -> Float32[Array, "b c"]:
if transpose:
return jnp.transpose(x)
else:
@@ -113,11 +141,11 @@ def test_nested_jaxtyped(typecheck, getkey):
def test_nested_nojaxtyped(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "b c"]):
def g(x: Float32[Array, "b c"]):
return h(x)
@typecheck
def h(x: f32[Array, "c b"]):
def h(x: Float32[Array, "c b"]):
return x
with pytest.raises(ParamError):
@@ -127,14 +155,14 @@ def test_nested_nojaxtyped(typecheck, getkey):
def test_isinstance(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "b c"]) -> f32[Array, " z"]:
def g(x: Float32[Array, "b c"]) -> Float32[Array, " z"]:
y = jnp.transpose(x)
assert isinstance(y, f32[Array, "c b"])
assert isinstance(y, Float32[Array, "c b"])
assert not isinstance(
y, f32[Array, "b z"]
y, Float32[Array, "b z"]
) # z left unbound as b!=c (unless x symmetric, which it isn't)
out = jr.normal(getkey(), (500,))
assert isinstance(out, f32[Array, "z"]) # z now bound
assert isinstance(out, Float32[Array, "z"]) # z now bound
return out
g(jr.normal(getkey(), (2, 3)))
@@ -143,7 +171,9 @@ def test_isinstance(typecheck, getkey):
def test_fixed(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "4 5 foo"], y: f32[Array, " foo"]) -> f32[Array, "4 5"]:
def g(
x: Float32[Array, "4 5 foo"], y: Float32[Array, " foo"]
) -> Float32[Array, "4 5"]:
return x @ y
a = jr.normal(getkey(), (4, 5, 2))
@@ -158,7 +188,7 @@ def test_fixed(typecheck, getkey):
def test_anonymous(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "foo _"], y: f32[Array, " _"]):
def g(x: Float32[Array, "foo _"], y: Float32[Array, " _"]):
pass
a = jr.normal(getkey(), (3, 4))
@@ -169,7 +199,11 @@ def test_anonymous(typecheck, getkey):
def test_named_variadic(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "*batch foo"], y: f32[Array, " *batch"], z: f32[Array, " foo"]):
def g(
x: Float32[Array, "*batch foo"],
y: Float32[Array, " *batch"],
z: Float32[Array, " foo"],
):
pass
c = jr.normal(getkey(), (5,))
@@ -189,7 +223,7 @@ def test_named_variadic(typecheck, getkey):
@jaxtyped
@typecheck
def h(x: f32[Array, " foo *batch"], y: f32[Array, " foo *batch bar"]):
def h(x: Float32[Array, " foo *batch"], y: Float32[Array, " foo *batch bar"]):
pass
a = jr.normal(getkey(), (4,))
@@ -205,7 +239,7 @@ def test_named_variadic(typecheck, getkey):
def test_anonymous_variadic(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "... foo"], y: f32[Array, " foo"]):
def g(x: Float32[Array, "... foo"], y: Float32[Array, " foo"]):
pass
a1 = jr.normal(getkey(), (5,))
@@ -227,7 +261,7 @@ def test_anonymous_variadic(typecheck, getkey):
def test_broadcast_fixed(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, "#4"]):
def g(x: Float32[Array, "#4"]):
pass
g(jr.normal(getkey(), (4,)))
@@ -240,7 +274,7 @@ def test_broadcast_fixed(typecheck, getkey):
def test_broadcast_named(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, " #foo"], y: f32[Array, " #foo"]):
def g(x: Float32[Array, " #foo"], y: Float32[Array, " #foo"]):
pass
a = jr.normal(getkey(), (3,))
@@ -264,7 +298,7 @@ def test_broadcast_named(typecheck, getkey):
def test_broadcast_variadic_named(typecheck, getkey):
@jaxtyped
@typecheck
def g(x: f32[Array, " *#foo"], y: f32[Array, " *#foo"]):
def g(x: Float32[Array, " *#foo"], y: Float32[Array, " *#foo"]):
pass
a = jr.normal(getkey(), (3,))
@@ -322,28 +356,28 @@ def test_broadcast_variadic_named(typecheck, getkey):
def test_no_commas():
with pytest.raises(ValueError):
f32[Array, "foo, bar"]
Float32[Array, "foo, bar"]
def test_symbolic(typecheck, getkey):
@jaxtyped
@typecheck
def make_slice(x: f32[Array, " dim"]) -> f32[Array, " dim-1"]:
def make_slice(x: Float32[Array, " dim"]) -> Float32[Array, " dim-1"]:
return x[1:]
@jaxtyped
@typecheck
def cat(x: f32[Array, " dim"]) -> f32[Array, " 2*dim"]:
def cat(x: Float32[Array, " dim"]) -> Float32[Array, " 2*dim"]:
return jnp.concatenate([x, x])
@jaxtyped
@typecheck
def bad_make_slice(x: f32[Array, " dim"]) -> f32[Array, " dim-1"]:
def bad_make_slice(x: Float32[Array, " dim"]) -> Float32[Array, " dim-1"]:
return x
@jaxtyped
@typecheck
def bad_cat(x: f32[Array, " dim"]) -> f32[Array, " 2*dim"]:
def bad_cat(x: Float32[Array, " dim"]) -> Float32[Array, " 2*dim"]:
return jnp.concatenate([x, x, x])
x = jr.normal(getkey(), (5,))
@@ -365,7 +399,7 @@ def test_symbolic(typecheck, getkey):
def test_incomplete_symbolic(typecheck, getkey):
@jaxtyped
@typecheck
def foo(x: f32[Array, " 2*dim"]):
def foo(x: Float32[Array, " 2*dim"]):
pass
x = jr.normal(getkey(), (4,))