# Thoughts on JAX vs Julia

**URL:** <https://discourse.julialang.org/t/thoughts-on-jax-vs-julia/86463>\
**Category:** Community\
**Tags:** jax\
**Created:** [August 28, 2022, 4:23pm UTC](https://discourse.julialang.org/t/thoughts-on-jax-vs-julia/86463 "2022-08-28T16:23:27Z")\
**Posts on this page:** 2\
**Page:** 2

<div class="post-metadata">

**Author:** ![czimm](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/czimm/32/20380_2.png) [@czimm](https://discourse.julialang.org/u/czimm)\
**Post date:** [September 13, 2022, 1:39am UTC](https://discourse.julialang.org/t/thoughts-on-jax-vs-julia/86463/21 "2022-09-13T01:39:14Z")

</div>

Hey Chris, where are you getting that `@pytime` macro? That seems like a huge nice-to-have and I can’t find it published anywhere.

---

<div class="post-metadata">

**Author:** ![jacobusmmsmit](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jacobusmmsmit/32/217669_2.png) [@jacobusmmsmit](https://discourse.julialang.org/u/jacobusmmsmit)\
**Post date:** [October 25, 2022, 11:06am UTC](https://discourse.julialang.org/t/thoughts-on-jax-vs-julia/86463/22 "2022-10-25T11:06:40Z")

</div>

Apologies for the delay in replying. My frustration with Julia in this regard essentially boils down to how imperative programming is slower than declarative programming because computers are imperative i.e. it’s not a frustration with Julia, but with computers themselves. JAX gets around this by being very limited in scope, and so it can make better optimisations from the same code due to being able to make more strict assumptions about what code will do.

While JAX isn’t fully declarative (whatever that means), I think `vmap` is a fantastic abstraction and I wish something with similar syntax existed in Julia. That said, JAX’s philosophy of function transformations + immutability works well with a lot of my code, but the times where it doesn’t (anything that doesn’t `vmap` nicely) I wish I could instead use Julia interleaved with blocks of JAX/XLA’s limitations in order to use its compiler.

On `vmap` and broadcasting. I much prefer the `vmap(f, (0, None))(x, y)` syntax over `f.(x, Ref(y))` as it clearly separates the verb from the object. Everyone loves abstraction.

The other alternatives in JAX are essentially just there to make compilation of `for` loops faster by providing a progressively less limiting function to allow you to do progressively more stuff that JAX wasn’t designed for. `lax.fori_loop` is implemented by `lax.scan` or `lax.while_loop`, and if `lax.fori_loop` won’t work then you have to use a regular `for` loop and suffer the compilation time of completely having it unrolled and each lime compiled separately. See [this implementation](https://github.com/jacobusmmsmit/multimixer/blob/master/multimixer/_src/backbone.py#L53) for an example when none of the loop primitives were possible, due to JAX not being able to `vmap` over an array of functions (also I need to remove that TODO).

[Previous page](https://discourse.julialang.org/t/thoughts-on-jax-vs-julia/86463.md?page=1)
