# PythonCall spends a lot of time showing stuff for JAX

**URL:** <https://discourse.julialang.org/t/pythoncall-spends-a-lot-of-time-showing-stuff-for-jax/121346>\
**Category:** Performance\
**Tags:** python, pythoncall, jax\
**Created:** [October 15, 2024, 6:59pm UTC](https://discourse.julialang.org/t/pythoncall-spends-a-lot-of-time-showing-stuff-for-jax/121346 "2024-10-15T18:59:40Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [October 15, 2024, 6:59pm UTC](https://discourse.julialang.org/t/pythoncall-spends-a-lot-of-time-showing-stuff-for-jax/121346/1 "2024-10-15T18:59:40Z")

</div>

Hi everyone, especially @cjdoris!

I’m trying to call JAX from Julia and I find that PythonCall.jl has a lot of overhead in a simple conversion case due to… printing? Has anyone run into the same issue?

Setup code:

```julia
using BenchmarkTools, CondaPkg, PythonCall
CondaPkg.add("numpy")
CondaPkg.add_pip("jax")
np = pyimport("numpy")
jnp = pyimport("jax.numpy")
x = rand(Float32, 1000);

```

Benchmark:

```julia
julia> @btime $(np.array)($x); # fast
  2.049 μs (17 allocations: 464 bytes)

julia> @btime $(jnp.array)($x); # slow
  2.400 ms (8386 allocations: 524.22 KiB)

```

Profiling:

```julia
julia> @profview for _ in 1:100; jnp.array(x); end

```

 ![image](https://global.discourse-cdn.com/julialang/original/3X/1/f/1fb3f4324bcbe372bb5e64db16e2ea68cc5c2e18.png)

---

<div class="post-metadata">

**Author:** ![stevengj](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/stevengj/32/71_2.png) [@stevengj](https://discourse.julialang.org/u/stevengj)\
**Post date:** [October 15, 2024, 7:17pm UTC](https://discourse.julialang.org/t/pythoncall-spends-a-lot-of-time-showing-stuff-for-jax/121346/2 "2024-10-15T19:17:57Z")

</div>

Weird, it seems like some conversion is going through `repr` rather than using the Python buffer interface. What if you call `jnp.asarray(x)`?

---

<div class="post-metadata">

**Author:** ![cjdoris](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cjdoris/32/213133_2.png) [@cjdoris](https://discourse.julialang.org/u/cjdoris)\
**Post date:** [October 15, 2024, 8:24pm UTC](https://discourse.julialang.org/t/pythoncall-spends-a-lot-of-time-showing-stuff-for-jax/121346/3 "2024-10-15T20:24:47Z")

</div>

Ok so it appears that `jnp.array(x)` is calling `str(x)` and `repr(x)` several times for some reason. I wonder if `jnp.array` tries a bunch of ways to convert `x` to an array, and the earlier tries involve throwing (then catching) an exception whose message includes `x`.

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [October 15, 2024, 8:59pm UTC](https://discourse.julialang.org/t/pythoncall-spends-a-lot-of-time-showing-stuff-for-jax/121346/4 "2024-10-15T20:59:43Z")

</div>

> [@stevengj](#):
>
> What if you call `jnp.asarray(x)`?

Same results unfortunately.

> [@cjdoris](#):
>
> I wonder if `jnp.array` tries a bunch of ways to convert `x` to an array, and the earlier tries involve throwing (then catching) an exception whose message includes `x`.

That’s probably what happens. By first using `np.array` followed by `jnp.array` I divide the overhead by 10:

```julia
julia> @btime $(np.array)($x); # fast
  3.073 μs (17 allocations: 464 bytes)

julia> @btime $(jnp.array)($(np.array(x)));
  105.150 μs (2 allocations: 32 bytes)

```

Of course the question is how much better this can get (it is a pretty crucial operation in [GitHub - gdalle/DifferentiationInterfaceJAX.jl](https://github.com/gdalle/DifferentiationInterfaceJAX.jl)).

EDIT: Perhaps I can get away with using only `np.array`, will report back.
