Forget, then remember: checkpointing the adjoint sweep on a GPU(nablatensor.com) |
Forget, then remember: checkpointing the adjoint sweep on a GPU(nablatensor.com) |
Fix is checkpointing (Griewank's 1992 "revolve," same idea modern LLM training uses under "gradient checkpointing"): keep ~38 bookmarks instead of 760 values, recompute what's needed when the reverse sweep asks for it. 2.9x, or 6.8x batching four markets per dispatch, bit-exact vs. the original.
One fun problem along the way: the compiler kept noticing the recomputation was identical to the original and merging them back together (undoing the whole point), so I had to XOR every bookmarked value with a runtime-zero constant just to stop it from proving that.