3 ms·
How is this different from JAX?
by roseandking 7y ago
How is this different from JAX?
- KenoFischer 7y agoJAX is a sophisticated and well implemented high level tracer, combined with a backend compiler that makes use of XLA. High level tracing has some fundamental draw backs in key features we want (efficient scalar AD for example), so the approach being advocated here is doing AD as a compiler transform on the original source language. Mike has some details in an earlier paper: https://arxiv.org/abs/1810.07951 https://arxiv.org/abs/1810.07951. As a shameless plug, you can plug (heh) Julia into XLA as well, e.g. for targeting TPUs (https://arxiv.org/abs/1810.09868 https://arxiv.org/abs/1810.09868). In fact, Julia and JAX are probably better at that than TensorFlow, because the TF execution model is a bit of a mismatch for what XLA expects.
- one-more-minute 7y agoThe really big problem with tracing, `@jit` annotations, and similar approaches, is that you're no longer running Python, but something almost-but-not-quite Python – modifying the language's semantics is necessary to get performance. In deep learning, tweaking some syntax and adding some annotations isn't a big deal, but for the use cases we're interested in we really don't want to rewrite every library to be AD compatible. Zygote is pretty unique (outside of the scientific computing world) in being able to take libraries that were written years before AD existed in Julia, and differentiate them correctly and efficiently. There are other, more subtle and technical, issues with those approaches, but that's really the Big Deal.