6 ms·
Not quite - it's more of a platform for users to write their own transformations on their PyTorch code. This can include things like doing operator fusion/lowe
by chillee 5y ago
Not quite - it's more of a platform for users to write their own transformations on their PyTorch code.
This can include things like doing operator fusion/lowering to a backend compiler, but can also include things like inserting profiling instrumentation (https://pytorch.org/tutorials/intermediate/fx_profiling_tutorial.html https://pytorch.org/tutorials/intermediate/fx_profiling_tuto...) or extracting intermediate features (https://github.com/pytorch/vision/releases/tag/v0.11.0 https://github.com/pytorch/vision/releases/tag/v0.11.0).
Basically, if you want a graph representation of a PyTorch module that's really easy to modify, use torch.fx :)
- mirker 5y agoBy “backend” I mean the compiler logic backing traced tensors in JAX. torch.fx seems different in that it’s primarily aimed at being a platform for users, rather than JAX which (as far as I know) hides the logic behind @jit annotations. Both trace eager code to a graph, which is then rewritten. JAX @jit is notably going eager -> graph -> XLA https://jax.readthedocs.io/en/latest/notebooks/How_JAX_primitives_work.html https://jax.readthedocs.io/en/latest/notebooks/How_JAX_primi... So (to me) it seems that they are similar backends/library primitives with different front-ends. There doesn’t seem to be a difference in representational power, since both hit a graph representation. The main exception I could see would be something like timers, which would perhaps require a graph-mode equivalent for JAX.
- chillee 5y agoWith 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.