{"slug": "trading-off-compute-for-memory-with-activation-checkpointing", "title": "Trading off compute for memory with activation checkpointing", "summary": "Jane Street ML Infra intern Arsh Koneru built an activation checkpointing scheme that, given a model spec and a memory budget, selects a set of nodes to drop from the save set to minimize extra compute time while staying within that budget. The project addresses PyTorch's existing mechanisms, including torch.compile's AOT Autograd min-cut partitioner, which drops cheap, fusible nodes from the save set, and the memory budget API. Activation checkpointing discards activations after the forward pass and recomputes them in the backward pass, trading compute for memory.", "body_md": "*The following is part of a series of posts about 2026 summer intern projects—for more, see [“What the interns have wrought, special jumbo 2026 edition”](https://blog.janestreet.com/wrought-2026/)*\n\nTo compute gradients in the backward pass of model training, we need the activations (outputs) of computations. By default, PyTorch saves these activations during the forward pass so that they can be reused during the backward pass. Once we compute the gradients, we can free up the memory for the activations, producing the typical memory profile seen above, with a peak at the end of the forward pass and a trough at the end of the backward pass.\n\n“Activation checkpointing” trades off some of this memory for compute. Instead of keeping\nan activation alive across the boundary, we discard it after we’ve used it in the forward\npass; when we need it again to compute gradients in the backward pass, we just re-run the\ncomputation that produced it. Doing this en masse would be a bad idea, since you’d end up\njust repeating the entire forward pass during the backward pass. But if you could avoid\nstoring some well-chosen nodes, you might find a better spot on the compute vs. memory\nfrontier. (The fact that this is called “checkpointing” is confusing, because in fact the\nwhole point is that we’re *not* saving the activation, but rather marking it as needing to\nbe recomputed later.)\n\nFor his project this summer, Arsh Koneru, a software engineering intern on our ML Infra team, was tasked with building a better activation checkpointing scheme. Given a model spec and a memory budget, he’d have to find a set of nodes to drop that would minimize extra compute time while staying within the budget.\n\n## Why not just use…?\n\nArsh’s first task was to understand existing tools for activation checkpointing already\nbuilt into PyTorch. In particular, when you call `torch.compile`, one of the phases is\ncalled AOT (Ahead-of-Time) Autograd. This computes a backward pass of your model in order\nto optimize over a joint graph both the forward and the backward. If there were no\nactivation checkpointing, every activation generated by the forward pass would survive in\nthe joint graph if it was directly needed by the backward pass. In the following\ncomputation, nodes **a** and **b** are saved from the forward pass but the backward pass\nneeds **b** and ***c*** (since they are inputs into the orange nodes, **u** and\n**v**). So, in the backward pass, we also recompute **c**.\n\nBy default when you run your model in eager mode, nothing needed for the backward pass is\never dropped from the save set (so we would’ve saved **b** and **c** in the diagram\nabove). But `torch.compile` applies an optimization. It drops some nodes from the save set\nusing the [min-cut\npartitioner](https://proceedings.mlsys.org/paper_files/paper/2023/file/8a27bb69950c0b46cdb36d10e5514cc8-Paper-mlsys2023.pdf),\nwhich improves runtime while also reducing memory. That may seem impossible, since we just\nsaid that activation checkpointing is a tradeoff between those two things; but there is a\ntrick here. When Inductor, torch.compile’s default backend, generates kernels for a graph,\nit fuses cheap ops together. And sometimes it is actually faster to recompute a forward op\nand fuse it into a backwards kernel than it is to store the forward op output in memory\nand load it into our kernel, because of the cost of transferring data into and out of\nmemory.\n\nNote in the figure that if both ops *g(x)* (from the forward pass) and *f(x)* (from the\nbackward pass) are memory-bound, then the fused kernel is actually faster, given the\nstipulated memory traffic.\n\nSo at a base level, PyTorch uses this min-cut optimization (with some restrictions and\nheuristics) to cut *some* nodes from the save set, namely those that are cheap and\nfusible.\n\nBut it also has another activation checkpointing mechanism given by the [memory budget\nAPI](https://pytorch.org/blog/activation-checkpointing-techniques/). What this does is\nrelax the min-cut rules in stages to propose more nodes to cut from the save set,\n[solving](https://github.com/pytorch/pytorch/blob/2622569e0df8fddefe23233c1087dbd6d07c3c21/torch/_functorch/partitioners.py#L3025)\na version of [the knapsack problem](https://en.wikipedia.org/wiki/Knapsack_problem) to\nchoose which ones.\n\nThis sounds like exactly the thing that Arsh was tasked with building, but in reality there were some major caveats with it:\n\n- We didn’t love the API. You have to specify the budget as a percentage of the default torch compile activations to save, which is hard for users to reason about. We want to be able to specify an absolute actual total GB memory budget.\n- It’s not monotonic: decreasing the budget from 20% to 10% can actually *increase* your\npeak memory usage (from extra memory used during the recomputation of the un-saved\nnodes).\n- It doesn’t actually account for the true memory peak, or the true amount of time taken to recompute a node.\n\nSo Arsh set out to build an alternative with all the features we wanted.\n\n## How to compute the true memory and recomputation cost of a node\n\nPyTorch actually also provides [Selective Activation Checkpointing\n(SAC)](https://pytorch.org/blog/activation-checkpointing-techniques/#:~:text=Policy%201%3A%20Not,CheckpointPolicy.PREFER_RECOMPUTE),\nwhich hands you each op in the graph and lets you decide whether to save or recompute it.\nBut the full graph structure is invisible through this API: you only get to see one op at\na time. This turns out not to be good enough.\n\nTo find the best ops to checkpoint, we need to see the entire graph, both to compute the true memory usage and the true recompute cost of each node.\n\n### Memory\n\nWe expect activation checkpointing to affect our memory usage like so:\n\nI.e., by eliminating a node from our save set in the forward pass, we expect the amount of memory not to climb as high during the forward pass, and therefore our peak should be lower.\n\nHowever, *it actually costs memory to recompute nodes in our graph*. So if we recompute\ntoo aggressively or in the wrong places, we can end up with a higher peak. For instance,\nif we need to allocate many or large intermediate tensors during our recomputation, our\npeak can move to the recompute itself, and we get something looking like this:\n\nThe checkpointing in this case both increases our peak memory and makes our training slower. Womp.\n\n### Recompute cost\n\nLikewise, the cost of recomputing a node is actually the cost of recomputing everything\nneeded to produce that node (or rather, everything that the backward pass would otherwise\nnot have needed). In the below example, nodes `a` and `b` originally do not need to be\ncomputed in the backward pass, since their values are only used to produce `c`, which is\nsaved.\n\nIf we then choose to checkpoint `c`, i.e., to not save it, then we must now recompute `a`\nand `b` in our backward pass as well. This makes the cost of not saving `c` equal to ```\na +\nb + c\n```\n.\n\n## The new activation checkpointing procedure\n\nBecause it needs to see the whole graph, Arsh’s algorithm for activation checkpointing requires an extra pass compiling the model. This “planning” pass is run once, cheaply, on a single GPU before training.\n\nThe planning pass compiles the model on fake tensors, which carry shapes and dtypes but no\ndata, and captures every graph it produces, plus the min-cut save set for each. (The\nmin-cut is already a good activation checkpointing baseline since it usually reduces both\nmemory and runtime. Starting there also reduces the search space, because nodes are fused,\nwhich makes planning faster.) While the graphs run, the planner also captures their\nexecution order, e.g. `A.fwd`, `B.fwd`, `B.bwd`, `A.bwd`, as those will come in handy\nlater.\n\nThen, given the graphs and their save sets, the planner searches for a subset of min-cut\nsaves to remove. (This new planner always does the same number or *fewer* saves, or rather\nfrees more memory, than the min-cut baseline.) In particular, it greedily drops the save\nwith the cheapest recompute cost per byte of peak memory freed. While the estimated peak\nis over the limit, for every output in the save set that can be dropped, the planner:\n\n1. Derives the new forward and backward graphs from the new save set\n2. Estimates the recompute cost of those new forward and backward graphs\n3. Simulates by how many bytes removing this save would reduce the *global* memory peak\n4. Removes the save having the best ratio, repeating until the limit is reached or no removal lowers the peak\n\nA greedy algorithm isn’t optimal and can get stuck in a local minimum. (Dropping one save could change every other candidate’s cost, and the algorithm won’t discover that because it never backtracks.) Something like beam search might be better, but this was a good start.\n\nTo estimate peak memory, the planner then traverses both the forward and backward graphs and tracks which activations are alive at every step. Inputs are alive for the whole graph, and outputs are alive from their creation until the end of the graph. Any intermediate transient memory is kept alive until its last use within the graph. Meanwhile, parameters, buffers, and optimizer state stay allocated for the whole run, so those are measured separately (also on fake tensors), along with a user-provided reserve for what the planner can’t see: fragmentation, NCCL buffers, CUDA overheads, etc.\n\nWhen there are multiple graphs stacked on top of each other, a tensor saved by graph A stays alive until A’s backward pass. Thus the true peak is the peak of that graph, plus everything the previous forward-pass graphs still have saved:\n\nTo estimate recompute time, the planner takes advantage of PyTorch’s [formulas that\nproduce the number of flops an operation takes given some input\nshape](https://github.com/pytorch/pytorch/blob/main/torch/utils/flop_counter.py). (This is\nused by torch’s own knapsack algorithm.) The planner also treats the total flops of a\ngraph as a proxy for how long it will take to recompute. In reality, not all flops are\nequal—it can depend on the datatype, on whether you’re using tensor cores, etc.—and not\nall compute-bound ops have a formula; *and* this completely ignores the time taken by\nmemory-bound kernels. Finally, we adjusted for some common ops in our models that are much\nslower than their flop formulas would suggest. After all this, the planner’s approximation\nended up being good enough for now.\n\n## Results\n\nAt training time, the plan is loaded and the joint graph is intercepted with a\n`custom_partitioner`. (Recall that the partitioner’s job is to take a joint graph and\nreturn the split forward and backward graphs.) Then, nodes are marked with `MUST_SAVE` and\n`MUST_RECOMPUTE`, min-cut is run again with these restriction tags, and that splits the\ngraph into forward and backward passes that match what was planned.\n\nOn the limited set of single-graph models that Arsh tested, auto-activation checkpointing (blue) beat the torch knapsack (red) and the op-level policies (orange) for most memory budgets. Points closer to the bottom-left of the graph indicate a faster activation checkpointing scheme for a given amount of peak memory.\n\nLooking at memory profiles shows that activation checkpointing moves the peak memory to the recomputation in the backward pass. And you can also see one of the advantages of planning against the full graph: no single recomputation spike dominates peak memory. Instead, the memory profile is smoothed out across time, which makes sense as a strategy for staying under a threshold.\n\nThere are lots of improvements still to be made—the algorithm can be optimized; it doesn’t\nwork on dynamic shapes or data-dependent ops (like `torch.nonzero`), and we can’t\nrecompute ops that have an RNG state; various memory estimates are crude approximations;\netc.—but Arsh’s work during the internship gives a promising baseline.\n\n## Looking forward to next summer…\n\n*If you’re interested in doing work like this, consider applying! You can find more\ndetails here: [Jane Street\nInternships](https://www.janestreet.com/join-jane-street/internships/). Applications for\nour 2027 internship are now open!*", "url": "https://wpnews.pro/news/trading-off-compute-for-memory-with-activation-checkpointing", "canonical_source": "https://blog.janestreet.com/trading-off-compute-for-memory-with-activation-checkpointing/", "published_at": "2026-09-30 00:00:00+00:00", "updated_at": "2026-09-30 23:18:24.441135+00:00", "lang": "en", "topics": ["machine-learning", "ai-infrastructure", "mlops", "developer-tools"], "entities": ["Jane Street", "Arsh Koneru", "PyTorch", "torch.compile", "Inductor", "ML Infra team", "AOT Autograd"], "also_reported_by": [], "alternates": {"html": "https://wpnews.pro/news/trading-off-compute-for-memory-with-activation-checkpointing", "markdown": "https://wpnews.pro/news/trading-off-compute-for-memory-with-activation-checkpointing.md", "text": "https://wpnews.pro/news/trading-off-compute-for-memory-with-activation-checkpointing.txt", "jsonld": "https://wpnews.pro/news/trading-off-compute-for-memory-with-activation-checkpointing.jsonld"}}