mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-09-11 12:21:38 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
81c56052e5 | ||
|
|
d911ebb99c | ||
|
|
f30b7d1546 | ||
|
|
4b3f834e12 |
@@ -26,7 +26,7 @@ jobs:
|
|||||||
run-tests:
|
run-tests:
|
||||||
strategy:
|
strategy:
|
||||||
matrix:
|
matrix:
|
||||||
python-version: [ 3.7, 3.8, 3.9 ]
|
python-version: [ 3.8, 3.9 ]
|
||||||
os: [ ubuntu-latest ]
|
os: [ ubuntu-latest ]
|
||||||
fail-fast: false
|
fail-fast: false
|
||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
|
|||||||
@@ -23,13 +23,13 @@ repos:
|
|||||||
hooks:
|
hooks:
|
||||||
- id: black
|
- id: black
|
||||||
- repo: https://github.com/nbQA-dev/nbQA
|
- repo: https://github.com/nbQA-dev/nbQA
|
||||||
rev: 1.2.3
|
rev: 1.6.3
|
||||||
hooks:
|
hooks:
|
||||||
- id: nbqa-black
|
- id: nbqa-black
|
||||||
- id: nbqa-isort
|
- id: nbqa-isort
|
||||||
- id: nbqa-flake8
|
- id: nbqa-flake8
|
||||||
- repo: https://github.com/PyCQA/isort
|
- repo: https://github.com/PyCQA/isort
|
||||||
rev: 5.10.1
|
rev: 5.12.0
|
||||||
hooks:
|
hooks:
|
||||||
- id: isort
|
- id: isort
|
||||||
- repo: https://github.com/pycqa/flake8
|
- repo: https://github.com/pycqa/flake8
|
||||||
|
|||||||
@@ -2,9 +2,10 @@
|
|||||||
|
|
||||||
Type annotations **and runtime checking** for:
|
Type annotations **and runtime checking** for:
|
||||||
|
|
||||||
1. shape and dtype of [JAX](https://github.com/google/jax) arrays;
|
1. shape and dtype of [JAX](https://github.com/google/jax) arrays; *(Now also supports PyTorch, NumPy, and TensorFlow!)*
|
||||||
2. [PyTrees](https://jax.readthedocs.io/en/latest/pytrees.html).
|
2. [PyTrees](https://jax.readthedocs.io/en/latest/pytrees.html).
|
||||||
|
|
||||||
|
|
||||||
**For example:**
|
**For example:**
|
||||||
```python
|
```python
|
||||||
from jaxtyping import Array, Float, PyTree
|
from jaxtyping import Array, Float, PyTree
|
||||||
|
|||||||
@@ -102,4 +102,4 @@ elif has_jax:
|
|||||||
|
|
||||||
del has_jax
|
del has_jax
|
||||||
|
|
||||||
__version__ = "0.2.11"
|
__version__ = "0.2.12"
|
||||||
|
|||||||
@@ -66,9 +66,11 @@ def _call_with_frames_removed(f, *args, **kwargs):
|
|||||||
|
|
||||||
def _optimized_cache_from_source(path, debug_override=None):
|
def _optimized_cache_from_source(path, debug_override=None):
|
||||||
# Version 2: change the position of the `@jaxtyped` decorator, so need a
|
# Version 2: change the position of the `@jaxtyped` decorator, so need a
|
||||||
# different name to avoid hitting old __pycache__
|
# different name to avoid hitting old __pycache__.
|
||||||
# Version 3: now also annotating classes.
|
# Version 3: now also annotating classes.
|
||||||
return cache_from_source(path, debug_override, optimization="jaxtyping3")
|
# Version 4: I'm honestly not sure, but bumping this fixed some kind of odd error.
|
||||||
|
# Maybe I changed something with hte classes part way through version 3?
|
||||||
|
return cache_from_source(path, debug_override, optimization="jaxtyping4")
|
||||||
|
|
||||||
|
|
||||||
def _dot_lookup(*elements):
|
def _dot_lookup(*elements):
|
||||||
@@ -111,7 +113,7 @@ class _JaxtypingTransformer(ast.NodeVisitor):
|
|||||||
args = [ast.Constant(None)]
|
args = [ast.Constant(None)]
|
||||||
else:
|
else:
|
||||||
args = [_dot_lookup(*self._typechecker)]
|
args = [_dot_lookup(*self._typechecker)]
|
||||||
node.decorator_list.append(ast.Call(func, args, keywords=[]))
|
node.decorator_list.insert(0, ast.Call(func, args, keywords=[]))
|
||||||
self._parents.append(node)
|
self._parents.append(node)
|
||||||
self.generic_visit(node)
|
self.generic_visit(node)
|
||||||
self._parents.pop()
|
self._parents.pop()
|
||||||
|
|||||||
@@ -17,6 +17,8 @@
|
|||||||
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
||||||
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
|
||||||
import equinox as eqx
|
import equinox as eqx
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
import pytest
|
import pytest
|
||||||
@@ -45,3 +47,16 @@ with pytest.raises(ParamError):
|
|||||||
M(1.0, jnp.array([1.0]))
|
M(1.0, jnp.array([1.0]))
|
||||||
with pytest.raises(ParamError):
|
with pytest.raises(ParamError):
|
||||||
M(1, jnp.array(1.0))
|
M(1, jnp.array(1.0))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class D:
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
D(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1, jnp.array(1.0))
|
||||||
|
|||||||
@@ -17,6 +17,8 @@
|
|||||||
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
||||||
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
|
||||||
import equinox as eqx
|
import equinox as eqx
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
import pytest
|
import pytest
|
||||||
@@ -45,3 +47,16 @@ with pytest.raises(ParamError):
|
|||||||
M(1.0, jnp.array([1.0]))
|
M(1.0, jnp.array([1.0]))
|
||||||
with pytest.raises(ParamError):
|
with pytest.raises(ParamError):
|
||||||
M(1, jnp.array(1.0))
|
M(1, jnp.array(1.0))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class D:
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
D(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1, jnp.array(1.0))
|
||||||
|
|||||||
@@ -17,6 +17,8 @@
|
|||||||
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
||||||
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
|
||||||
import equinox as eqx
|
import equinox as eqx
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
import pytest
|
import pytest
|
||||||
@@ -45,3 +47,16 @@ with pytest.raises(ParamError):
|
|||||||
M(1.0, jnp.array([1.0]))
|
M(1.0, jnp.array([1.0]))
|
||||||
with pytest.raises(ParamError):
|
with pytest.raises(ParamError):
|
||||||
M(1, jnp.array(1.0))
|
M(1, jnp.array(1.0))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class D:
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
D(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1, jnp.array(1.0))
|
||||||
|
|||||||
+15
-2
@@ -1,13 +1,26 @@
|
|||||||
|
import abc
|
||||||
|
|
||||||
from jaxtyping import jaxtyped
|
from jaxtyping import jaxtyped
|
||||||
|
|
||||||
|
|
||||||
class M:
|
class M(metaclass=abc.ABCMeta):
|
||||||
@jaxtyped
|
@jaxtyped
|
||||||
@classmethod
|
@classmethod
|
||||||
def f(cls):
|
def f(cls):
|
||||||
return 3
|
return 3
|
||||||
|
|
||||||
|
@jaxtyped
|
||||||
|
@abc.abstractmethod
|
||||||
|
def g(self):
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
# Check that the @jaxtyped decorator doesn't blat the __get__ of @classmethod
|
# Check that the @jaxtyped decorator doesn't blat the __get__ of @classmethod
|
||||||
def test_decorator():
|
def test_classmethod():
|
||||||
assert M.f() == 3
|
assert M.f() == 3
|
||||||
|
|
||||||
|
|
||||||
|
# Check that the @jaxtyped decorator doesn't blat the __isabstractmethod__ of
|
||||||
|
# @abstractmethod
|
||||||
|
def test_abstractmethod():
|
||||||
|
assert M.g.__isabstractmethod__
|
||||||
|
|||||||
Reference in New Issue
Block a user