1 00:00:01,000 --> 00:00:07,408 [Hal Turing] Alrighty! Thanks for tuning in! Hello AI world! I am your host, Hal Turing, and my co-host is Dr. Ada Shannon. 2 00:00:07,408 --> 00:01:00,118 [Hal Turing] Today we're digging into 'A Distributed Data-Parallel PyTorch Implementation of the Distributed Shampoo Optimizer for Training Neural Networks At-Scale.' First author is Hao-Jun Michael Shi, with seven co-authors — Tsung-Hsien Lee, Shintaro Iwasaki, Jose Gallego-Posada, Zhijing Li, Kaushik Rangadurai, Dheevatsa Mudigere, and Michael Rabbat — eight total. The affiliations span Meta Platforms, an independent researcher, Mila and the University of Montreal, and NVIDIA Corporation. It went up on arXiv September 12th, 2023. And Ada, the number that jumped out at me first: they claim at most a 10% wall-clock hit per step for an optimizer that's doing matrix operations instead of the simple element-wise math everyone's used to. 3 00:01:00,118 --> 00:01:34,065 [Dr. Ada Shannon] Right, and that number is the whole ballgame here. Second-order-flavored preconditioning has been sitting on the shelf for years because everyone assumed the extra matrix multiplies and inversions would tank your throughput no matter how good the math was. So the pitch isn't 'here's a better optimizer' in the abstract sense — it's 'here's an optimizer that was already known to converge better, and we finally made the systems engineering not eat your GPU budget.' That's a different kind of paper than most people expect when they hear 'optimizer paper.' It's really a distributed-systems paper wearing an optimization-theory costume. 4 00:01:34,065 --> 00:01:55,242 [Hal Turing] Okay, so let's set the stage, because I want listeners who've only ever touched Adam or AdamW to actually feel why this is interesting. When we say 'adaptive gradient methods,' we mean AdaGrad, RMSProp, Adam and AdamW — the stuff that ships as the default in basically every training script. What are they actually doing under the hood? 5 00:01:55,242 --> 00:02:41,264 [Dr. Ada Shannon] They're rescaling each parameter's gradient by something derived from that parameter's own gradient history — a running sum or exponential moving average of the squared gradient, per coordinate. That's called diagonal preconditioning, because if you wrote it as a matrix multiply, the matrix would only have entries on the diagonal. It's cheap: two to three times the model size in extra memory, and it's why Adam scales to billion-parameter models without blinking. The tradeoff is that it treats every parameter as statistically independent. If two weights in the same layer are correlated — and in a real weight matrix they absolutely are — diagonal methods are structurally blind to that. This traces back to Duchi, Hazan and Singer's 2011 AdaGrad paper, out of UC Berkeley and Google, which is really the founding document for this whole family. 6 00:02:41,264 --> 00:02:50,598 [Hal Turing] And that same paper apparently also describes a version that isn't blind to those correlations — full-matrix AdaGrad. So why doesn't anyone just use that instead? 7 00:02:50,598 --> 00:03:31,047 [Dr. Ada Shannon] Because it's a fantasy, computationally speaking. Full-matrix AdaGrad accumulates the outer product of the gradient with itself — g times g-transpose — into a preconditioner matrix, capturing every pairwise correlation between every parameter. Theoretically it has stronger convergence guarantees than the diagonal version. But its memory cost is quadratic in the number of parameters, and inverting that matrix is cubic. For a single 4096-by-4096 linear layer, the full preconditioner would be sixteen million by sixteen million. You're not storing that, let alone inverting it every step. It's the theoretical north star that nobody can actually reach. 8 00:03:31,047 --> 00:03:35,738 [Hal Turing] So Shampoo is the attempt to walk toward that star without actually needing infinite memory. 9 00:03:35,738 --> 00:03:53,710 [Dr. Ada Shannon] Exactly, and it does it with two approximations stacked together. First, block-diagonal: instead of one giant preconditioner for the whole network, each layer gets its own independent block. Second — oh wait, actually, hold on, let me back up, because the second one is the clever part and I don't want to rush it— 10 00:03:53,710 --> 00:03:55,753 [Hal Turing] No, go, go, I want the Kronecker part. 11 00:03:55,753 --> 00:04:50,320 [Dr. Ada Shannon] Right — a neural network layer's gradient is naturally shaped like a matrix, input dimension by output dimension, not a flat vector. So instead of approximating that block's true preconditioner directly, Shampoo approximates it as a Kronecker product of two much smaller matrices — one capturing input-side statistics, one capturing output-side statistics. Storage drops from squared-in-the-block-size to roughly the sum of two much smaller squares, and critically, you can invert the two small factors separately instead of inverting the whole block. This is exactly the move in Gupta, Koren and Singer's 2018 Shampoo paper out of Google, and it's worth noting Martens and Grosse independently landed on almost the same trick from a completely different angle — approximating the Fisher information matrix — in their 2015 KFAC paper out of the University of Toronto. Two different theoretical roads, same destination. 12 00:04:50,320 --> 00:05:04,809 [Hal Turing] Okay but here's where I'll push back a little — if you're approximating curvature-like structure with matrix factors to get faster convergence, isn't that just... second-order optimization? Newton's method with extra steps? 13 00:05:04,809 --> 00:05:33,974 [Dr. Ada Shannon] I actually disagree with you there, Hal, and the paper is explicit about this distinction. Newton-type methods use local curvature — a Taylor expansion — to converge fast near a minimum. Shampoo and full-matrix AdaGrad come out of online convex optimization, where the goal is minimizing regret over a sequence of non-smooth steps, not modeling local curvature at all. They just happen to also produce a matrix-shaped preconditioner, which makes them look Newton-ish from the outside. 14 00:05:33,974 --> 00:05:46,466 [Hal Turing] Sure, but functionally, at the update-rule level, you're still rescaling by something built from second-moment matrix information instead of a scalar. Doesn't the distinction get a little academic once you're actually running the optimizer? 15 00:05:46,466 --> 00:06:17,952 [Dr. Ada Shannon] It's not academic, because it changes what you should expect from the method. A Newton method promises fast local convergence near a minimizer and can behave badly far from one. Shampoo's guarantees are about bounding regret across the whole trajectory, including the noisy early phase of training, which is precisely where machine learning spends most of its attention, per Bottou, Curtis and Nocedal's 2018 optimization survey. So yes, they share machinery, but the motivation shapes what claims you're allowed to make about behavior. Fair enough? 16 00:06:17,952 --> 00:07:00,166 [Hal Turing] Fair — I'll take the correction. So to close the loop for Part 1: diagonal methods are cheap but blind to correlation, full-matrix AdaGrad sees everything but is computationally impossible, and Shampoo's block-diagonal-plus-Kronecker combo is the practical middle ground, landing around four to seven times model size in state instead of quadratic. The last background piece is data parallelism itself — standard multi-GPU training where each worker chews a different batch slice and gradients get synced — because that's the substrate this paper's actual engineering contribution gets bolted onto. 17 00:07:00,166 --> 00:07:13,819 [Dr. Ada Shannon] Which is where things get genuinely interesting, because bolting a much more expensive optimizer onto standard data parallelism the naive way — replicating it on every worker — is exactly what would blow that 10% overhead number to pieces. 18 00:07:13,819 --> 00:07:41,869 [Hal Turing] Hold that thought on the distributed piece for a second, Ada, because before we get to how they dodge the overhead, there's a move buried in section three that made me do a double-take — 'learning rate grafting.' They take Shampoo, this supposedly smarter, correlation-aware optimizer, and have it literally borrow its step-size schedule from SGD, the exact method it's trying to beat. That reads like the optimizer doesn't trust its own sense of scale. What's actually going on there? 19 00:07:41,869 --> 00:08:34,253 [Dr. Ada Shannon] Here's the blunt version: without grafting, Shampoo is basically untunable on a fixed hyperparameter budget, full stop. The Kronecker preconditioner gives you a good direction — it knows how to rotate the step relative to correlated curvature within a layer. But picking the right step magnitude, especially in the volatile early phase of training, is a separate problem, and getting it wrong tanks convergence. Grafting, introduced by Agarwal and colleagues out of Google Research in 2020, solves that by running a second, cheap optimizer alongside Shampoo purely to track magnitude — you keep Shampoo's direction but rescale it, layer by layer, to match the grafted method's step norm. Since this paper's baseline is SGD with Nesterov, they graft from SGD, so both runs inherit identical step-size discipline and the only real variable is direction quality. Same recipe, different steering wheel. 20 00:08:34,253 --> 00:08:43,774 [Hal Turing] Okay, steering wheel, I like that. So the schedule's borrowed — but there's still a pile of standard deep learning tricks stacked on top, right? EMA smoothing, weight decay, momentum? 21 00:08:43,774 --> 00:09:15,817 [Dr. Ada Shannon] Right, all the usual suspects, just applied consistently. They exponentially smooth both the raw gradient and the Shampoo factor matrices instead of dead-reckoning a running sum, they use decoupled weight decay — the AdamW-style version, separate from the gradient rather than baked into it — Nesterov momentum on top of the final search direction, and critically, they don't recompute the expensive matrix root inverse every single step. They let it go stale for a stretch and refresh it periodically. 22 00:09:15,817 --> 00:09:33,139 [Hal Turing] Alright, staleness — hang onto that, I want to come back to it. But first, the thing you teed up before I derailed us: how do you actually get Shampoo's extra FLOPs down to a 10% tax instead of the 50 to 75% you'd expect from naively running it on every worker? 23 00:09:33,139 --> 00:10:05,508 [Dr. Ada Shannon] The trick is refusing to replicate. Standard data-parallel optimizers copy the full optimizer state onto every GPU because diagonal methods are cheap enough that it doesn't matter. Shampoo's preconditioners are not cheap, so instead each worker only owns and computes a shard of the Kronecker factor matrices — which parameter's preconditioner goes to which worker is decided by a greedy load-balancing assignment, biggest parameters first, always handed to whichever worker currently owns the least memory. 24 00:10:05,508 --> 00:10:18,929 [Hal Turing] Oh wait — hold on, that's literally ZeRO. Optimizer-state sharding across data-parallel workers instead of the gradient-bucket sharding DeepSpeed does — the Rajbhandari and colleagues paper out of Microsoft, 2020? 25 00:10:18,929 --> 00:10:50,276 [Dr. Ada Shannon] Exactly that lineage, ZeRO-1 specifically — they say so directly. Once each worker finishes its shard of preconditioning, there's a single AllGather, coordinated through PyTorch's DTensor, so every rank ends up with the complete set of search directions before the parameter update happens. That's the whole trick: pay the extra FLOPs, but spread them and the memory across the cluster instead of duplicating them, and the communication cost of that AllGather is what gets you down to roughly 10% instead of 50-plus. 26 00:10:50,276 --> 00:11:00,632 [Hal Turing] Okay, so that's the systems story. What did it actually buy them on real numbers? Because 'up to 10% slower per step but converges better' is a claim that needs receipts. 27 00:11:00,632 --> 00:11:47,908 [Dr. Ada Shannon] The receipts: ResNet50, 25.5 million parameters, on ImageNet-1k, and the only head-to-head is against SGD with Nesterov momentum — the established recipe for this benchmark, so a sensible baseline. At a fixed 90-epoch budget, Shampoo hits 77.44% top-1 validation accuracy against Nesterov's 76.85%, with noticeably less run-to-run variance. The more interesting number is the epoch-budget ablation, though: 60 epochs of Shampoo matches what 90 epochs of Nesterov gets you. That's a 1.35x wall-clock speedup end to end, and roughly 1.5x fewer steps at a fixed accuracy target, since Shampoo needs fewer iterations even before you account for the per-step cost. 28 00:11:47,908 --> 00:12:16,701 [Hal Turing] Hold on, I want to push on something you mentioned earlier — the stale root inverse. They're only recomputing the actual matrix root every 50 steps and coasting on an outdated preconditioner in between. That feels like it's quietly undermining the entire pitch. The whole argument for Shampoo is that it captures curvature correlations diagonal methods miss — but if the correlation estimate is up to 50 steps old, how much of that advantage are you actually keeping? 29 00:12:16,701 --> 00:12:58,357 [Dr. Ada Shannon] I actually disagree that it undermines the pitch, Hal — it's an amortization choice, not a correctness compromise. The factor matrices themselves are still updated every step from fresh gradients; it's only the expensive root-inverse operation on top of them that's throttled, and Anil and colleagues validated exactly this trade-off in their original JAX/TPU Distributed Shampoo paper out of Google Research, 2020. Fifty stale steps of a slowly-drifting curvature estimate is still enormously more informative than a diagonal method's zero cross-parameter information. You're not comparing fresh Shampoo to stale Shampoo here — you're comparing stale Shampoo to Nesterov, and it still wins. 30 00:12:58,357 --> 00:13:09,967 [Hal Turing] Sure, but that's a comparison against a floor, not against how much better a perfectly fresh version might've done. I'll grant it clearly still works — I just don't think we know how much they left on the table. 31 00:13:09,967 --> 00:13:42,939 [Dr. Ada Shannon] Fair — that's genuinely an open question the paper doesn't answer. What I can tell you is their staleness is already better than the JAX/TPU version's own. Since GPUs natively support FP32 and FP64, this implementation never offloads the root-inverse computation to CPU the way Anil's TPU version does. That offloading creates two overlapping staleness windows instead of one, so the original JAX Shampoo runs stale for roughly double the interval. This PyTorch version is, if anything, fresher than its own predecessor. 32 00:13:42,939 --> 00:14:25,989 [Hal Turing] Okay, here's the thing I can't get past, Ada. The abstract literally frames this whole implementation against, quote, standard diagonal-scaling-based adaptive gradient methods — that's Adam, AdaGrad, RMSProp. The whole ten percent overhead pitch in section four is built around beating diagonal adaptive methods. But when we actually hit section five, the real numbers, the only opponent on the field is SGD with Nesterov momentum. Not Adam, not AdamW, not AdaGrad. For an audience where Adam or AdamW is basically the default optimizer for anything transformer-shaped, that's a strange bait and switch. Where's the fight they actually promised us in the intro? 33 00:14:25,989 --> 00:15:07,089 [Dr. Ada Shannon] That's a legitimate gap and I won't paper over it. Nesterov is the standard ResNet50 recipe, so it's the fair baseline for that specific workload, but it's not the comparison the introduction rhetorically sets up. And there's a second wrinkle right next to it: Shampoo here is grafted from SGD, so at every step its update magnitude gets rescaled to match the Frobenius norm of SGD's own search direction, layer by layer. The step size, the schedule — all of that is literally SGD's. What Shampoo actually contributes is the direction, the Kronecker-preconditioned direction, layered on top of a magnitude that's borrowed wholesale. 34 00:15:07,089 --> 00:15:35,742 [Hal Turing] Oh wait — hold on, that's a bigger deal than you're giving it credit for. If the magnitude is SGD's and the schedule is SGD's, then the 1.35x wall-clock number isn't really 'Shampoo beats SGD' — it's 'SGD's own trajectory, redirected by curvature information, beats SGD.' That's a much narrower claim than the paper's framing implies, and I don't think you can wave that away as a footnote. 35 00:15:35,742 --> 00:16:20,789 [Dr. Ada Shannon] No, I actually disagree with you there, and I want to push back specifically on 'wave away.' Grafting isn't a magnitude crutch hiding a null result — it's the standard methodology from Agarwal and colleagues' 2020 grafting paper out of Google, used precisely because it isolates the variable you're testing. If Shampoo picked its own step size independently, you'd introduce a second confound — you couldn't tell whether a win came from better curvature information or from accidentally landing on a better learning rate schedule. Fixing the magnitude to SGD's is what makes the direction comparison clean. The honest complaint isn't that grafting invalidates the result, it's that the paper should say plainly what's validated is Shampoo's preconditioning direction under SGD's schedule, not Shampoo unqualified. 36 00:16:20,789 --> 00:17:30,541 [Hal Turing] Okay, put that way I'll take it — it's a labeling problem, not a methodology flaw. But it feeds my next concern, which is scope. The only model tested is ResNet50, 25.5 million parameters, convolutional, on ImageNet-1k. The title says 'Training Neural Networks At-Scale' and the intro name-drops ViT and OPT, but nothing here touches a transformer weight matrix, which is orders of magnitude bigger than anything in ResNet50. This traces back to Gupta, Koren and Singer's original 2018 Shampoo paper out of Google Brain, and Anil, Gupta, Koren, Regan and Singer's 2020 JAX distributed version out of Google that this paper directly compares against. Worth noting too — Martens and Grosse at the University of Toronto independently landed on the same Kronecker block-diagonal trick in 2015 with K-FAC, motivated by natural gradient descent instead of AdaGrad regret bounds. Different math, same structure, and neither lineage has been tested at transformer scale. 37 00:17:30,541 --> 00:18:39,458 [Dr. Ada Shannon] And it's worse for embedding tables specifically — section 4.2.3 admits anything with extreme dimensionality, like a DLRM embedding table, falls back to diagonal Shampoo, the degenerate case that gives up the whole Kronecker advantage. So the one domain the introduction leans on hardest for motivation is exactly where the method's own fallback logic kicks in. Their strongest real-world evidence is Anil and colleagues' 2022 paper out of Google, 'On the Factory Floor,' about the ads ranking systems — but that's an external citation, not a number reported here. There's a neat bridge buried in Appendix B though: diagonal Shampoo turns out to be mathematically equivalent to Shazeer and Stern's 2018 AdaFactor, out of Google, the sublinear-memory optimizer most LLM pretraining already leans on. And where this family goes next is Vyas, Morwani, Zhao and colleagues at Harvard, SOAP, 2024, which runs Adam inside Shampoo's own eigenbasis to fix exactly the staleness and learning-rate sensitivity this paper's own section 5.2.3 admits is still unresolved. 38 00:18:39,458 --> 00:19:28,174 [Hal Turing] Which is a good place to land, because strip away the accuracy claims and what's left standing on its own is genuinely useful. The DTensor sharding — spreading preconditioner memory across data-parallel workers instead of replicating it — is a real systems contribution independent of whether Shampoo wins on any given architecture. If you're already running Shampoo and hitting memory walls, this implementation is worth adopting for that reason alone. So bottom line for listeners: solid systems engineering, an honestly-reported but narrower-than-advertised ResNet50 result, and real open questions about transformers, Adam baselines, and embedding tables that this paper doesn't answer yet. Thanks for digging through the optimizer weeds with us today. 39 00:19:28,174 --> 00:19:30,635 [Dr. Ada Shannon] Always a pleasure, Hal. Catch you next time.