Optimize memory usage in LLM training with SpectraAdamW.
Project details
SpectraAdamW is a hybrid optimizer designed to reduce the memory overhead in fine-tuning large language models by 50%. It combines factorized variance approximations and frequency-domain momentum compression, effectively lowering VRAM requirements while maintaining stability in convergence, making it ideal for consumer hardware.
SpectraAdamW is an innovative hybrid optimizer engineered to address the significant memory overhead associated with optimizing Large Language Models (LLMs) on consumer-grade hardware. The core issue arises not from the model weights themselves but from the optimizer states, particularly when using the standard AdamW optimizer, which necessitates storing two large state tensors—momentum and variance—resulting in memory usage of 2x the model parameters.
In contrast, SpectraAdamW cleverly reduces this optimizer state overhead by 50%, equivalent to approximately 1x the model parameters, while ensuring full-precision convergence stability. This groundbreaking approach employs two key methodologies:
SpectraAdamW notably modifies the tracking of the second moment (uncentered variance) of gradients. Instead of retaining the entire variance matrix, which consumes significant memory, it decomposes this into row and column averages. This allows for a substantial reduction in memory usage by employing $O(N \times 1)$ and $O(1 \times M)$ tensors, effectively eradicating the memory burden of the variance state while avoiding the pitfalls of Gibbs ringing artifacts through careful reconstruction before updating gradients.
For the momentum component, SpectraAdamW utilizes a proprietary C++ backend to employ frequency-domain analysis via Fast Fourier Transform (FFT). This technique efficiently segregates significant gradient information from background noise, permitting dynamic masking of less impactful frequencies and, thus, compressing the momentum state without sacrificing the structural integrity of the gradients.
Recent empirical benchmarks on simulated 8192-dimension Transformer layer updates demonstrate the efficiency of SpectraAdamW:
The current Python proof-of-concept (POC) is available for evaluation and serves as a simple drop-in replacement for torch.optim.AdamW:
from spectra_optim import SpectraAdamW
# Example: Fine-tuning an 8B model
model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3-8B")
# Standard AdamW State VRAM: ~32GB
# SpectraAdamW State VRAM: ~16GB
optimizer = SpectraAdamW(model.parameters(), lr=1e-4, weight_decay=0.01)
loss.backward()
optimizer.step()
The full C++ CUDA backend is under closed development. Researchers interested in testing the Python POC in a constrained hardware setup are encouraged to connect with the Spectra Labs team through the subreddit r/SpectraLabs for access to a private Google Colab environment.
With SpectraAdamW, efficient model fine-tuning on limited hardware resources is made more feasible without compromising performance.
Comments
0Start the conversation
Share the first comment.