5 ms·
With this kind of stuff, I think the devil is in the details, but the principle is similar, yes. Specifically, the level of abstraction at which you're tracing,
by chillee 5y ago
With this kind of stuff, I think the devil is in the details, but the principle is similar, yes. Specifically, the level of abstraction at which you're tracing, as well as what types of programs you can express/not express.
For example, FX is extremely unopinionated about what it can trace, and the trace itself is extremely customizable. For example, if a subfunction/module has control flow (i.e. untraceable), it's easy to mark it as a "leaf" in FX's tracer, while that concept doesn't really make in sense in Jax's tracing system.
Another example of a difference is that Jax traces out into its own IR called a jaxpr, while FX is explicitly a Python => Python translation. This has some some upsides and some downsides - for example, you can insert arbitrary Python functions into your FX graph (breakpoints, print statements, etc.), while jaxprs don't allow that.
Is this a good thing? Well, if your main goal is to lower to XLA, definitely not lol. But for FX it works quite well.
TL;DR: The general principles of doing graph capture are similar, but the details matter, and the details end up being quite different.