This repository was archived by the owner on Sep 9, 2026. It is now read-only.
File tree Expand file tree Collapse file tree
computation_backends/jax_backend Expand file tree Collapse file tree Original file line number Diff line number Diff line change 99 AnyTensor ,
1010 AudioTensor ,
1111 ImageTensor ,
12- JaxArray ,
1312 NdArray ,
1413 VideoTensor ,
1514)
1918if jax_available :
2019 import jax .numpy as jnp
2120
21+ from docarray .typing import JaxArray
22+
2223
2324@pytest .fixture ()
2425def batch ():
Original file line number Diff line number Diff line change 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
1016def test_cosine_sim_jax ():
Original file line number Diff line number Diff line change 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
715def test_top_k_descending_false ():
Original file line number Diff line number Diff line change 1- import jax .numpy as jnp
21import numpy as np
32import pytest
4- from jax ._src .core import InconclusiveDimensionOperation
53from pydantic import schema_json_of
64from pydantic .tools import parse_obj_as
75
86from 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
1217def test_proto_tensor ():
You can’t perform that action at this time.
0 commit comments