GWT: Scalable Optimizer State Compression for Large Language Model Training

arXiv:2501.07237 · cs.LG, cs.AI · Submitted 2025-01-13 · Read on arXiv

Listen

Radio episode about this paper

Transcript

Introduction to the show: ident: AI Radio. Generated commentary on the latest Artificial Intelligence papers.

Tom: Next we'll be talking about the paper "GWT: Scalable Optimizer State Compression for Large Language Model Training".

Jane: The paper was written by Ziqing Wen, Ping Luo, Jiahuan Wang, Kun Yuan, Dongsheng Li et al. from National Key Laboratory of Parallel and Distributed Computing, National University of Defense Technology and Center for Machine Learning Research, Peking University.

Tom: Stay tuned as we take you through the paper and discuss its implications.

Summary: Tom: So, let's get into the core idea of GWT, which is summarized in the abstract and introduction. It addresses how stateful optimizers like Adam impose massive memory overhead during LLM training.

Jane: You mentioned that Figure one shows Adam consumes twice the model memory; GWT aims to fix that by compressing gradients using wavelets instead of just relying on rank reduction, which is a much more nuanced approach.

Lu: The paper explains that GWT projects gradients into wavelet subspaces, essentially capturing only the essential update information while allowing for massive compression.

Meng: It's interesting to see they aren't limited to low-rank assumptions like LoRA or GaLore, so Meng wonders how that affects the practical applicability across different architectures.

Lalam: The fact that this compression is tied to preserving essential update information means the potential for generating meaningful and coherent AI outputs should also be preserved.

Tom: And GWT achieves remarkable results, specifically showing up to a seventy-nine percent reduction in optimizer memory when applying it to LLaMA models during pre-training.

Jane: That level of compression, coupled with a one point nine times training speedup on the C4 dataset, is what makes this paper so immediately exciting for engineers and researchers alike.

Lu: It’s a significant shift from merely optimizing existing methods to inventing a new pathway for efficient learning dynamics.

Meng: We're looking at substantial improvements in hardware utilization here, especially when scaling up training runs.

Lalam: This reduction in overhead should allow for more ambitious training runs that lead to better cultural outcomes and understanding of humanity.

Improvements: Tom: We've seen the memory savings, but how much better is GWT than existing low-rank methods like GaLore or APOLLO?

Jane: The paper argues that by utilizing the Wavelet Transform—capturing local information features and suppressing high-frequency noise—it offers a distinct advantage over methods that just discard information outside a specific subspace.

Lu: It’s not just a heuristic approach; they provide rigorous mathematical proofs showing that under certain conditions, the Haar low-pass approximation outperforms any global rank-r approximation in terms of the Frobenius norm.

Meng: That theoretical grounding is crucial, because it suggests this method isn't just a lucky hack; it provides stability and efficiency at scale.

Lalam: Stability is a word that resonates with me, because chaotic or noisy gradients can lead to unpredictable AI behavior, so ensuring stable learning is vital for reliable deployment.

Tom: The paper also mentions its robustness to long sequence lengths, which is a huge practical concern in LLM design where context window size matters.

Lu: They show that while GaLore degrades noticeably as the sequence length increases, GWT maintains stable performance across these varying contexts.

Meng: From an engineering view, this suggests that we can deploy GWT in scenarios with long context windows without having to drastically change our hardware or training parameters.

Lalam: That robustness translates into reliability for users, ensuring that the AI we build is consistent whether it's answering a quick question or processing a massive document.

Conclusion: Tom: We have established that GWT saves memory and is theoretically sound, but what does the data tell us about actual performance?

Jane: The results in Table II are incredibly strong, showing GWT consistently achieving lower validation perplexity (PPL) than full-rank Adam across various LLaMA models.

Lu: The theoretical backing of Theorem five confirms that this level of efficiency is not just a theoretical possibility but an achievable reality in the over-parameterized regime.

Meng: We also see performance parity with advanced methods like GaLore and APOLLO, which is a massive win for practical deployment at scale.

Lalam: It’s encouraging to see that GWT doesn't compromise the quality of the model; we can get both efficiency and high-quality knowledge extraction.

Tom: And it doesn's just about pre-training; Table III shows its speed advantage, with a one point nine times increase in token throughput for LLaMA 3B using GWT-two.

Lu: The fact that this works across different levels of the wavelet transform also adds to the flexibility and potential of this approach.

Meng: This speedup translates directly into reduced cloud costs, which is a major consideration for AI startups and large research labs alike.

Lalam: It seems like GWT is enabling a future where we can train larger, smarter AI more efficiently, leading us toward greater collective understanding.

Wrap-up: Tom: Before we wrap up this segment, I want to summarize our discussion on "GWT: Scalable Optimizer State Compression for Large Language Model Training."

Jane: We’ve seen that GWT addresses the massive memory burden of Adam optimizers by compressing gradients using wavelet transforms.

Lu: This method offers a powerful combination of mathematical rigor and practical efficiency, proving that traditional low-rank assumptions are not the only way to achieve optimal performance in LLM training.

Meng: It’s a scalable, robust solution that provides significant hardware advantages and can be applied to diverse optimizers beyond just Adam.

Lalam: We hope this work leads to an era where we can train AI models of incredible complexity without being constrained by the sheer physical limits of memory and processing power.

Tom: And I think that's a powerful message for everyone listening today. We've seen how "GWT: Scalable Optimizer State Compression for Large Language Model Training" is fundamentally changing how we approach large-scale AI development.

Jane: It’s been a truly insightful conversation, and I hope this discussion inspires some further thoughts on the world to come.

Tom: Thank you all for joining us, and we'll be back with more research next time!

National Key Laboratory of Parallel and Distributed Computing, National University of Defense Technology, Changsha, 410073, China · Center for Machine Learning Research, Peking University, Beijing 100871, China

cs.LG, cs.AI

Submitted: 2025-01-13

Updated: 2026-08-25

Code: https://github.com/tatsu-lab/stanford_alpaca

Importance score: 89/100

The gist: GWT (Gradient Wavelet Transform) is presented as a scalable method designed for "Optimizer State Compression for Large Language Model Training," aiming to improve memory efficiency during LLM

Key concepts

GWT
GWT is an approach that compresses gradients using wavelet subspaces instead of relying on low-rank assumptions. It captures essential update information, allowing for massive compression of stateful optimizers like Adam during LLM training.
Wavelet Transform
The Wavelet Transform is used by GWT to capture local information features and suppress high-frequency noise in gradients. This offers a distinct advantage over methods that simply discard data outside a specific subspace.

Terminology

Summary

GWT (Gradient Wavelet Transform) is presented as a scalable method designed for Optimizer State Compression for Large Language Model Training, aiming to improve memory efficiency during LLM training while maintaining high performance.

Compression Mechanism and Trade-offs:

The efficacy of the compression technique is analyzed concerning the wavelet level, denoted by l. The experimental findings indicate that the main impact of l seems to be on memory usage and throughput. Specifically, there is a clear trade-off: Higher wavelet levels improve memory usage, but this may come at the cost of slightly higher validation perplexity and slower training speed. Furthermore, the paper raises an important research question regarding the compression coefficients themselves: whether the approximation coefficients are not crucial when using wavelets to compress gradients in LLM training, but rather the detail coefficients.

Quantitative Performance Analysis (Memory and Throughput):

The performance trade-offs are quantitatively summarized across multiple model sizes (e.g., 60M, 130M, 350M). Table XII provides a detailed comparison of ESTIMATED OPTIMIZER MEMORY USAGE AND TOKEN THROUGHPUT. This table demonstrates a consistent pattern: HIGHER GWT LEVELS LEAD TO SLOWER TRAINING TOKEN THROUGHPUT BUT SMALLER MEMORY USAGE. For instance, comparing the memory usage across different GWT levels (GWT-1 through GWT-5) shows a continuous reduction in required memory, coupled with a corresponding decrease in token throughput.

Integration of Modular Learning Rate Strategies:

A key component of the GWT implementation is the adoption of a modular learning rate strategy. This strategy involves partitioning the overall learning rate such that the Attention and MLP modules are updated using a scaled learning rate of lr × α, while the rest of the model uses the base learning rate lr. This approach is noted to be commonly adopted in other memory-efficient training algorithms, including LoRA, GaLore, Fira, and APOLLO.

The paper conducted an ablation experiment comparing GWT's performance against Adam when using this modular learning rate scheme. The results showed that Adam significantly benefits from the modular learning rate strategy, achieving clearly better performance than with a single global learning rate. Crucially, even when GWT is configured with this enhanced setting, GWT achieves a final validation perplexity (PPL) close to that of Full-Adam.

Conclusion and Implications:

These findings lead to several significant conclusions regarding the architecture and training dynamics of LLMs. The superior performance achieved by Adam using modular learning rates suggests that the modular learning rate strategy plays a key role in memory-efficient methods outperforming full-rank Adam. Furthermore, the results suggest that Attention and MLP modules may require smaller learning rates than other parts of the model. This leads to the overarching question posed by the research: Do different modules in LLMs inherently demand distinct learning rates? And could explicitly integrating modular learning rate strategies into Adam lead to further gains?

Improvements for AI systems

The integration of the Gradient Wavelet Transform (GWT) offers several critical, high-impact improvements across large-scale AI systems:

  • Improvement: Dramatic reduction in memory overhead required for stateful optimizers (like Adam). GWT achieves up to 79% reduction in optimizer memory usage when projecting gradients using high-level wavelet transforms (l=3 or higher).

  • System Capability: This allows large-scale LLM training pipelines to utilize significantly fewer GPUs (e.g., requiring substantially less than the 9 additional A100 80GB GPUs previously needed for a 175B parameter model) while maintaining the same computational fidelity.

  • Improvement: GWT overcomes the inherent performance degradation associated with traditional low-rank methods (LoRA, GaLore, APOLLO). The method captures essential update information by leveraging wavelet subspaces rather than discarding it.

  • System Capability: The improved system can achieve final validation scores (e.g., PPL) that are comparable to or superior to full-rank updates and significantly better than state-of-the-art low-rank baselines, ensuring high model quality even under severe memory constraints.

  • Improvement: GWT facilitates more efficient information compression during the gradient update process, leading to faster convergence rates compared to full-rank methods.

  • System Capability: The improved system exhibits higher token throughput, with reported speeds up to 1.9x faster than standard methods (e.g., achieving over 0.532K tokens/s per GPU). This enables the processing and training of massive datasets in a fraction of the time, accelerating research cycles and deployment readiness.

  • Improvement: GWT is designed to be optimizer-agnostic, meaning it functions effectively when integrated with various optimization protocols (Adam, Adam-mini, MUON) without requiring specialized modifications to the existing training framework. Furthermore, its performance remains stable even when handling long sequence lengths.

  • System Capability: The system can be seamlessly integrated into any existing deep learning pipeline and is reliably applicable to complex scenarios involving long-context windows (e.g., 1024 tokens) without experiencing the performance degradation seen in other compression techniques like GaLore.


In summary, the GWT-enhanced AI system transforms LLM training from a resource bottleneck into a scalable, high-fidelity process that is significantly faster and cheaper to execute.

Sources

Related papers