It sounds like JAX is necessarily storing your whole calculation in memory, so it will necessarily use more memory for automatic differentiation of heavily iterative calculations, while other implementations of backward-mode automatic differentiation can instead restart your calculation from checkpoints to avoid storing the whole thing. This could be an advantage of several orders of magnitude for some calculations: using twice the CPU or GPU time in exchange for one thousandth or one ten-thousandth of the memory.