Functorch: Jax-like composable function transforms for PyTorch
github.com
github.com
For example, an unjitted function’s error can be easier to debug in jax than in anything else, while the same error in the jitted function can be harder to debug in jax than in anything else. But it’s similar with vmap, pmap, grad, and the rest of the transformations. Debugging gets nuts quickly.
But I don’t think there’s any way around that with these kinds of transformations in python, is there?