Skip to content
This repository was archived by the owner on Sep 9, 2026. It is now read-only.

Commit 92e32ec

Browse files
committed
fix: is_jax_available added for missing tests
Signed-off-by: agaraman0 <agaraman0@gmail.com>
1 parent 27e3c13 commit 92e32ec

4 files changed

Lines changed: 32 additions & 12 deletions

File tree

tests/units/array/stack/test_array_stacked_jax.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,6 @@
99
AnyTensor,
1010
AudioTensor,
1111
ImageTensor,
12-
JaxArray,
1312
NdArray,
1413
VideoTensor,
1514
)
@@ -19,6 +18,8 @@
1918
if jax_available:
2019
import jax.numpy as jnp
2120

21+
from docarray.typing import JaxArray
22+
2223

2324
@pytest.fixture()
2425
def batch():

tests/units/computation_backends/jax_backend/test_metrics.py

Lines changed: 11 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,16 @@
1-
import jax
2-
import jax.numpy as jnp
1+
from docarray.utils._internal.misc import is_jax_available
32

4-
from docarray.computation.jax_backend import JaxCompBackend
5-
from docarray.typing import JaxArray
3+
jax_available = is_jax_available()
4+
if jax_available:
5+
import jax
6+
import jax.numpy as jnp
67

7-
metrics = JaxCompBackend.Metrics
8+
from docarray.computation.jax_backend import JaxCompBackend
9+
from docarray.typing import JaxArray
10+
11+
metrics = JaxCompBackend.Metrics
12+
else:
13+
metrics = None
814

915

1016
def test_cosine_sim_jax():

tests/units/computation_backends/jax_backend/test_retrieval.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,15 @@
1-
import jax.numpy as jnp
1+
from docarray.utils._internal.misc import is_jax_available
22

3-
from docarray.computation.jax_backend import JaxCompBackend
4-
from docarray.typing import JaxArray
3+
jax_available = is_jax_available()
4+
if jax_available:
5+
import jax.numpy as jnp
6+
7+
from docarray.computation.jax_backend import JaxCompBackend
8+
from docarray.typing import JaxArray
9+
10+
metrics = JaxCompBackend.Metrics
11+
else:
12+
metrics = None
513

614

715
def test_top_k_descending_false():

tests/units/typing/tensor/test_jax_array.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,17 @@
1-
import jax.numpy as jnp
21
import numpy as np
32
import pytest
4-
from jax._src.core import InconclusiveDimensionOperation
53
from pydantic import schema_json_of
64
from pydantic.tools import parse_obj_as
75

86
from docarray.base_doc.io.json import orjson_dumps
9-
from docarray.typing import JaxArray
7+
from docarray.utils._internal.misc import is_jax_available
8+
9+
jax_available = is_jax_available()
10+
if jax_available:
11+
import jax.numpy as jnp
12+
from jax._src.core import InconclusiveDimensionOperation
13+
14+
from docarray.typing import JaxArray
1015

1116

1217
def test_proto_tensor():

0 commit comments

Comments
 (0)