mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-09-10 12:14:04 +08:00
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:
@@ -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 }}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -102,4 +102,4 @@ elif has_jax:
|
||||
|
||||
del has_jax
|
||||
|
||||
__version__ = "0.2.11"
|
||||
__version__ = "0.2.12"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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__
|
||||
|
||||
Reference in New Issue
Block a user