# Using Julia autodiff code from python with JAX

**URL:** <https://discourse.julialang.org/t/using-julia-autodiff-code-from-python-with-jax/132062>\
**Category:** New to Julia\
**Tags:** python, autodiff\
**Created:** [September 2, 2025, 6:19pm UTC](https://discourse.julialang.org/t/using-julia-autodiff-code-from-python-with-jax/132062 "2025-09-02T18:19:36Z")\
**Posts on this page:** 1\
**Showing post:** 2

<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:** [September 2, 2025, 6:40pm UTC](https://discourse.julialang.org/t/using-julia-autodiff-code-from-python-with-jax/132062/2 "2025-09-02T18:40:12Z")

</div>

I think your best bet combines Reactant.jl with Enzyme.jl. I’m not sure exactly how to combine both to do what you want, but @wsmoses will know.  
You may also be interested in past discussion on the topic:

> [@Calling Python function JIT compiled with JAX from Julia without overhead](https://discourse.julialang.org/t/calling-python-function-jit-compiled-with-jax-from-julia-without-overhead/123552):
>
> It could be very useful to call JIT compiled JAX functions from Python libraries from Julia, and in turn, to be able for Julia libraries to call user-provided JAX functions. This is possible today via PythonCall, e.g.: using PythonCall jax = pyimport("jax") numpy = pyimport("numpy") # Define a simple test function func\_str = """ def f(x): return jax.numpy.sin(x) + jax.numpy.cos(x) """ # Create namespace and define function namespace = pydic…

---

_[View the full topic](https://discourse.julialang.org/t/using-julia-autodiff-code-from-python-with-jax/132062)._
