PolarGrad: A Class of Matrix-Gradient Optimizers from a Unifying Preconditioning Perspective
August 10, 2026 · View on GitHub
PolarGrad (Polar Gradient methods; Lau et al., 2025) is a class of matrix-gradient optimizers based on the concept of gradient-anisotropy preconditioning in optimization. It has close relation to Muon (Jordan et al., 2024) and stochastic spectral descent (SSD; Carlson et al., 2015a, 2015b). In addition to being an optimizer for matrix parameters in neural networks, PolarGrad can also be used as a preconditioned matrix optimization algorithm for matrix optimization problems such as matrix regression and low-rank matrix factorization/completion.
The main differences between PolarGrad and Muon/SSD are:
- PolarGrad uses the QDWH (Nakatsukasa et al., 2010) or ZOLO-PD (Nakatsukasa and Freund, 2016) algorithm to compute the polar decomposition of the gradient matrix, while Muon uses the Newton-Schulz iteration to compute the polar decomposition (see the section below for further details). The NS iteration is a matrix iterative polynomial method that computes the polar decomposition of a matrix by iteratively applying a polynomial to the matrix. However, it requires tuning of the coefficients of the polynomial, which can be challenging in practice. PolarGrad also include the nuclear norm (the dual norm of the spectral norm) scaling of the update matrix, which is not present in Muon. The inclusion of such term is necessary for the convergence of optimizers based on polar decomposition for strongly convex and Lipschitz smooth problems with deterministic gradients, as shown in the convergence analysis and the matrix quadratic regression example of PolarGrad (Lau et al., 2025).
- Following the concurrent work of Amsel et al. (2025), PolarGrad also includes the Polar Express method to compute the polar decomposition of a matrix, which uses polynomial approximations of the sign function to compute the polar decomposition, rather than rational approximations of the sign function used in the QDWH and ZOLO-PD algorithms, hence avoiding the use of QR decompositions and involving only matrix-matrix products (in half-precision arithmetic). Its implementation is directly taken from the paper. Yet, it has not been heavily tested in experiments yet.
- While SSD also includes the nuclear norm scaling, PolarGrad uses more advanced numerical linear algebra algorithms for polar decomposition than the randomized SVD algorithm used in SSD, namely the QDWH and ZOLO-PD algorithms.
Overview
This repository provides implementations of PolarGrad in PyTorch utilizing two more advanced numerical linear algebra algorithms for polar decomposition than the Newton-Schulz (NS) iteration:
- The QWDH algorithm (Nakatsukasa et al., 2010; see here and here for implementation in JAX)
- The ZOLO-PD algorithm (Nakatsukasa and Freund, 2016; see here for the authors' MATLAB implementation)
These two algorithms, unlike the NS iteration, do not require tuning of the coefficients of the matrix iterative polynomial, and they are more numerically stable (Nakatsukasa and Higham, 2012; Nakatsukasa and Freund, 2016). Hence, they are more suitable for matrix parameters of different sizes and potentially ill-conditioned initializations, making them a better candidate and optimizers based on polar decomposition like PolarGrad and Muon (Jordan et al., 2024) a drop-in replacement of other adaptive gradient optimizers such as Adam(W). Currently, the QWDH algorithm is particularly more efficient for large matrices, while ZOLO-PD is designed for small to medium-sized matrices. Note that both of these algorithms involve QR decompositions, which might not be efficient for GPUs and half-precision arithmetic. To addresss such issue, we also include the Polar Express method in Amsel et al. (2025) to compute the polar decomposition of a matrix, which uses polynomial approximations of the sign function to compute the polar decomposition, rather than rational approximations of the sign function used in the QDWH and ZOLO-PD algorithms, hence avoiding the use of QR decompositions and involving only matrix-matrix products (in half-precision arithmetic).
In particular, with the assist of ChatGPT, we translated these implementations in JAX and MATLAB to PyTorch. Currently, limited by the QR decomposition implementation in PyTorch, mixed precisions such as bfloat16 are not yet supported. Notice that the current implementation is not optimized for speed and parallelization, although we have also provided a DDP implementation polar_grad_ddp.py, following the implementation of Muon. The three main files are:
-
polar.py: includes the functionpolarwhich mimics the JAXjax.scipy.linalg.polarfunction, which computes the polar decomposition of a matrix using four possible numerical algorithms.i.
method=qdwh: uses the QDWH algorithm (Nakatsukasa et al., 2010) to compute the polar decomposition of a matrix. This is suitable for large matrices and is more numerically stable than the Newton-Schulz iteration.ii.
method=zolo-pd: uses the ZOLO-PD algorithm (Nakatsukasa and Freund, 2016) to compute the polar decomposition of a matrix. This is suitable for small to medium-sized matrices and is also more numerically stable than the Newton-Schulz iteration.iii.
method=ns: uses the Newton-Schulz (NS) iteration to compute the polar decomposition of a matrix. This might require tuning of the coefficients of the matrix iterative polynomial for different model and layer sizes, which can be challenging in practice. This is the same method used in the Muon optimizer (Jordan et al., 2024), and is adopted from its GitHub repository.iv.
method=precond_ns: uses the preconditioned Newton-Schulz iteration in Lewis et al. (2022) to compute the polar decomposition of a matrix. This is potentially an improved variant of the NS iteration with the need of coefficient tuning, but might still suffer from the stability issue of the NS iteration. We include this method for completeness, but is not heavily tested and not used in the experiments in the paper.v.
method=polar_express: uses the Polar Express method in Amsel et al. (2025) to compute the polar decomposition of a matrix. -
polar_grad.py: includes thetorch.optim.OptimizerclassPolarGradwhich implements the PolarGrad optimizer based on the above four numerical polar decomposition algorithms of the gradient matrix.- The argument
polar_firstspecifies whether polar-first momentum is used; default isFalsewhich is similar to the implementation of Muon (Jordan et al., 2024). - The argument
methodspecifies which polar decomposition algorithm to use, and can be one of the following:qdwh(cf.qdwh.pyadopted from its JAX implementationjax.lax.linalg.qdwh),zolo-pd(cf.zolopd.pyadopted from its MATLAB implementation),nsorprecond_ns(cf.newton_schulz.pyadopted from Muon's GitHub repository). The default isqdwh, which is suitable for large matrices. - The argument
inner_stepsspecifies the number of (inner) steps for either the QDWH algorithm or the NS iteration. The other two algorithms (ZOLO-PD and preconditioned NS) do not require this argument. The default is2. - The arguments
a,bandcspecify the coefficients of the matrix iterative polynomial for the NS iteration, which are used only whenmethod='ns'. The default values are the same as those in Muon, which are suitable for most cases for hidden layers. However, they can be tuned for different model and layer sizes if necessary.
The optimizer can be used as follows:
optimizer = PolarGrad(model.parameters(), lr=1e-3, weight_decay=0., momentum=0.9, polar_first=False, method='qdwh', inner_steps=2) - The argument
-
polar_grad_ddp.py: includes thetorch.optim.OptimizerclassPolarGradwhich implements the PolarGrad optimizer based on the above four numerical polar decomposition algorithms of the gradient matrix withtorch.distributed, following the implementation in Muon's GitHub repository.
Installation of Required Libraries
Install PyTorch (nightly) accodring to the instructions at https://pytorch.org/get-started/locally/, e.g., for Linux and CUDA 12.6:
pip install --pre torch torchvision torchaudio --index-url https://download.pytorch.org/whl/nightly/cu126
Then, install some auxiliary libraries:
pip install -U numpy matplotlib tqdm fire SciencePlots
For correct LaTeX rendering in matplotlib, you might also need to have a LaTeX distribution installed, such as TeX Live (MacTeX) or MikTeX, or disable the LaTeX rendering in matplotlib by setting rcParams['text.usetex'] = False and changing some of the plot labels in the code.
Usage
For small-scale experiments which can be run with CPU, you can run the following commands to test the PolarGrad optimizer on different matrix optimization problems. The --seed argument is used to set the random seed for reproducibility.
-
Matrix quadratic regression (a strongly convex problem with deterministic gradient):
# PolarGrad python mat_quad_reg.py --steps=4000 --seed=42 # PolarGradM python mat_quad_reg_mom.py --steps=4000 --seed=42 -
Matrix logistic regression (a strongly convex problem with stochastic gradient):
# PolarSGD python mat_log_reg.py --steps=1500 --seed=42 # PolarSGDM python mat_log_reg_mom.py --steps=1500 --seed=42 -
Low-rank matrix completion (a non-convex problem with deterministic gradient):
# PolarGrad python low_rank_mat_comp.py --steps=1000 --seed=42 # PolarGradM python low_rank_mat_comp_mom.py --steps=200 --seed=42
We will update the repository with examples and experiments for language model pre-training soon.
Timing, memory, and GPU benchmarks for Sections 6.1--6.3
The three main experiment scripts have an opt-in benchmark mode. The benchmark
uses a fresh copy of each workload and omits the condition-number and
nuclear-norm diagnostics from the timed region. This is important because those
diagnostics can cost substantially more than the optimizer step. A throwaway
warmup run absorbs lazy CUDA initialization and torch.compile overhead, while
the measured run starts again from the same seed and model initialization.
Commit the benchmark code before collecting final numbers so the JSON metadata
contains an immutable Git commit and reports a clean worktree.
To run only the timing and memory measurements on a CUDA GPU:
python mat_quad_reg.py --device=cuda --benchmark_only=True --benchmark_repeats=3
python mat_log_reg.py --device=cuda --benchmark_only=True --benchmark_repeats=3
python low_rank_mat_comp.py --device=cuda --benchmark_only=True --benchmark_repeats=3
Use --benchmark=True instead of --benchmark_only=True to generate the
original iteration-based figures first and then run the independent benchmark.
By default, the benchmark measures the same number of steps as the experiment.
For a short validation run, pass (for example) --benchmark_steps=100. The
options --benchmark_warmup, --benchmark_repeats, and --results_dir control
the warmup length, number of fresh repeats, and output directory.
Each script writes a JSON record, per-repeat CSV, and summary CSV. The records
include total wall-clock time, time per step, steps per second, forward/backward
and optimizer device time, optimizer-state size, baseline CUDA allocation, and
peak allocated/reserved CUDA memory. The estimated temporary workspace is the
incremental peak minus the final optimizer-state tensor size; it should be
treated as an approximation because caching allocators can reuse storage.
To keep event instrumentation from dominating these small workloads, component
times are sampled every ten steps and their sampling interval is stored in the
output. device_activity_pct is the estimated fraction of
wall time covered by CUDA-event intervals. It is a timing-based activity proxy,
not SM occupancy or achieved FLOP efficiency; the saved system metadata should
always accompany reported numbers. CPU runs report timing but leave CUDA-memory
fields at zero.
The polar-oracle microbenchmark compares fixed inner-iteration budgets while also reporting orthogonality, reconstruction, and optional exact-direction errors:
python benchmark_polar_oracles.py \
--device=cuda \
--shapes=500x100,1000x100,500x5,250x5,1024x1024,4096x1024 \
--methods=qdwh,zolo-pd,ns,polar_express \
--spectra=gaussian,ill_conditioned,rank_deficient \
--inner-steps=2,5 \
--condition-number=1e6 \
--compute-reference \
--repeats=3
Add --profile to save PyTorch Chrome traces for kernel-level inspection. For
publication-quality hardware utilization (for example, SM and DRAM throughput),
run the same microbenchmark under NVIDIA Nsight Compute or Nsight Systems and
report those profiler metrics separately from device_activity_pct.
Reproducibility corrections made with this benchmark update:
- The experiment scripts now pass the reported inner-iteration budgets
explicitly: two QDWH steps and five Newton--Schulz steps. Previously,
Muon_polarsilently used its default of five steps for QDWH. - Matrix logistic regression now generates
Cby thresholding standard Gaussian samples at0.5, as stated in the experimental details. - Low-rank matrix completion now applies the observation mask to the residual in both the joint optimizer and AltGD objectives. Previously, the mask was generated but used only in the denominator.
Citation
If you find this repository useful for your research, please consider citing our paper using the BibTeX entry below:
@article{lau2025polargrad,
title={\textsc{PolarGrad}: A Class of Matrix-Gradient Optimizers from a Unifying Preconditioning Perspective},
author={Lau, Tim Tsz-Kit and Qi Long and Weijie Su},
year={2025},
journal={arXiv preprint arXiv:2505.21799}
}
References
-
Lau, Tim Tsz-Kit, Qi Long, and Weijie Su. PolarGrad: A class of matrix-gradient optimizers from a unifying preconditioning perspective. arXiv preprint arXiv:2505.21799, 2025.
-
Jordan, Keller, Yuchen Jin, Vlado Boza, Jiacheng You, Franz Cesista, Laker Newhouse, and Jeremy Bernstein. Muon: An optimizer for hidden layers in neural networks. 2024.
-
Carlson, David, Volkan Cevher, and Lawrence Carin. Stochastic spectral descent for restricted Boltzmann machines. In Proceedings of the International Conference on Artificial Intelligence and Statistics (AISTATS), 2015a.
-
Carlson, David E., Edo Collins, Ya-Ping Hsieh, Lawrence Carin, and Volkan Cevher. Preconditioned spectral descent for deep learning. In Advances in Neural Information Processing Systems (NeurIPS), 2015b.
-
Nakatsukasa, Yuji, Zhaojun Bai, and François Gygi. Optimizing Halley's iteration for computing the matrix polar decomposition. SIAM Journal on Matrix Analysis and Applications, 31(5):2700-2720, 2010.
-
Nakatsukasa, Yuji and Roland W. Freund. Computing fundamental matrix decompositions accurately via the matrix sign function in two iterations: The power of Zolotarev's functions. SIAM Review, 58(3):461-493, 2016.
-
Nakatsukasa, Yuji, and Nicholas J. Higham. Backward stability of iterations for computing the polar decomposition. SIAM Journal on Matrix Analysis and Applications, 33(2):460-479, 2012.
-
Lewis, Adam G. M., Jackson Beall, Martin Ganahl, Markus Hauru, Shrestha Basu Mallick, and Guifre Vidal. Large-scale distributed linear algebra with tensor processing units. Proceedings of the National Academy of Sciences, 119(33):e2122762119, 2022.
-
Amsel, Noah, David Persson, Christopher Musco, and Robert M. Gower. The Polar Express: Optimal matrix sign methods and their application to the Muon algorithm. arXiv preprint arXiv:2505.16932, 2025.