mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-08-21 11:16:29 +08:00
* Add a test for generators * Remove output annotations from decorators Also guarded torch imports for better compatibility with requirements.txt * Add flag to the main meta class to skip the typecheck * Return to the old solution * Make async tests work * Minor adjustments/fixing typos * Correct Python path for new tests * Remove some jax-dependent code * Implement equality for MetaArrays * Make all Dim variations frozen dataclasses * Shorten AbstractArray methods * Final touches * Removing get_origin use * Update tests with @jaxtyp
26 lines
497 B
Python
26 lines
497 B
Python
import cloudpickle
|
|
import numpy as np
|
|
|
|
|
|
try:
|
|
import torch
|
|
except ImportError:
|
|
torch = None
|
|
|
|
from jaxtyping import AbstractArray, Array, Shaped
|
|
|
|
|
|
def test_pickle():
|
|
x = cloudpickle.dumps(Shaped[Array, ""])
|
|
cloudpickle.loads(x)
|
|
|
|
y = cloudpickle.dumps(AbstractArray)
|
|
cloudpickle.loads(y)
|
|
|
|
z = cloudpickle.dumps(Shaped[np.ndarray, ""])
|
|
cloudpickle.loads(z)
|
|
|
|
if torch is not None:
|
|
w = cloudpickle.dumps(Shaped[torch.Tensor, ""])
|
|
cloudpickle.loads(w)
|