Fixes for some new failures. (Where did they come from?) (#65)

* Fixes for some new failures. (Where did they come from?)

* Fixed isort?
This commit is contained in:
Patrick Kidger
2023-02-16 10:08:57 -08:00
committed by GitHub
parent d911ebb99c
commit 926470732b
6 changed files with 24 additions and 14 deletions
+1 -1
View File
@@ -26,7 +26,7 @@ jobs:
run-tests:
strategy:
matrix:
python-version: [ 3.7, 3.8, 3.9 ]
python-version: [ 3.8, 3.9 ]
os: [ ubuntu-latest ]
fail-fast: false
runs-on: ${{ matrix.os }}
+2 -2
View File
@@ -23,13 +23,13 @@ repos:
hooks:
- id: black
- repo: https://github.com/nbQA-dev/nbQA
rev: 1.2.3
rev: 1.6.3
hooks:
- id: nbqa-black
- id: nbqa-isort
- id: nbqa-flake8
- repo: https://github.com/PyCQA/isort
rev: 5.10.1
rev: 5.12.0
hooks:
- id: isort
- repo: https://github.com/pycqa/flake8
+1 -1
View File
@@ -102,4 +102,4 @@ elif has_jax:
del has_jax
__version__ = "0.2.11"
__version__ = "0.2.12"
+1 -6
View File
@@ -31,12 +31,7 @@ class _Jaxtyped:
self.fn = fn
def __get__(self, instance, owner):
if instance is None:
# Don't create a new _Jaxtyped object in this case. Otherwise anything
# assigned to methods (e.g. `__isabstractmethod__`) just gets swallowed.
return self
else:
return ft.wraps(self.fn)(_Jaxtyped(self.fn.__get__(instance, owner)))
return ft.wraps(self.fn)(_Jaxtyped(self.fn.__get__(instance, owner)))
def __call__(self, *args, **kwargs):
try:
+4 -2
View File
@@ -66,9 +66,11 @@ def _call_with_frames_removed(f, *args, **kwargs):
def _optimized_cache_from_source(path, debug_override=None):
# 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.
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):
+15 -2
View File
@@ -1,13 +1,26 @@
import abc
from jaxtyping import jaxtyped
class M:
class M(metaclass=abc.ABCMeta):
@jaxtyped
@classmethod
def f(cls):
return 3
@jaxtyped
@abc.abstractmethod
def g(self):
...
# Check that the @jaxtyped decorator doesn't blat the __get__ of @classmethod
def test_decorator():
def test_classmethod():
assert M.f() == 3
# Check that the @jaxtyped decorator doesn't blat the __isabstractmethod__ of
# @abstractmethod
def test_abstractmethod():
assert M.g.__isabstractmethod__