OKLS optimizer - #265
Conversation
Signed-off-by: mkhona <mkhona@nvidia.com>
Signed-off-by: mkhona <mkhona@nvidia.com>
Signed-off-by: mkhona <mkhona@nvidia.com>
Signed-off-by: mkhona <mkhona@nvidia.com>
Signed-off-by: mkhona <mkhona@nvidia.com>
Signed-off-by: mkhona <mkhona@nvidia.com>
Signed-off-by: mkhona <mkhona@nvidia.com>
Signed-off-by: mkhona <mkhona@nvidia.com>
| if ridge_eps < 0.0: | ||
| raise ValueError(f"Invalid ridge epsilon: {ridge_eps}") |
There was a problem hiding this comment.
Zero epsilon produces NaN state
When ridge_eps=0 and the first gradient is all zero, initialization computes sqrt(rows / 0) and multiplies the zero covariance by infinity, causing NaN factors, inverse roots, and parameter values on the first step.
| if ridge_eps < 0.0: | |
| raise ValueError(f"Invalid ridge epsilon: {ridge_eps}") | |
| if ridge_eps <= 0.0: | |
| raise ValueError(f"Invalid ridge epsilon: {ridge_eps}") |
| grad_right_preconditioned = grad @ inverse_root_right | ||
| factor_left.lerp_(grad_right_preconditioned @ grad_right_preconditioned.T / cols, 1 - shampoo_beta) | ||
| factor_left.copy_((factor_left + factor_left.T) / 2.0) | ||
| factor_left.diagonal().add_(ridge_eps) | ||
|
|
||
| grad_left_preconditioned = inverse_root_left @ grad | ||
| factor_right.lerp_(grad_left_preconditioned.T @ grad_left_preconditioned / rows, 1 - shampoo_beta) |
There was a problem hiding this comment.
Factor matmuls inherit global precision
The covariance and final preconditioning matmuls execute outside fp32_matmul_precision, so a process-wide medium setting silently reduces their precision and makes persistent optimizer factors depend on unrelated global configuration.
Knowledge Base Used: SOAP: Shampoo-style Preconditioning
Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!
Greptile SummaryAdds the OKLS optimizer and its scaled-CANS inverse-root implementation.
Confidence Score: 3/5The PR is not yet safe to merge because zero epsilon can corrupt optimizer state with NaNs and configured matmul precision is not applied throughout the OKLS update. The current implementation still permits a zero epsilon that makes zero-gradient initialization evaluate zero times infinity, while factor and final preconditioning matmuls remain controlled by unrelated process-wide precision. Files Needing Attention: emerging_optimizers/soap/okls.py Important Files Changed
Flowchart%%{init: {'theme': 'neutral'}}%%
flowchart LR
G[Gradient] --> I[Initialize or update Kronecker factors]
I --> C[Scaled CANS inverse roots]
G --> M[Nesterov momentum]
C --> P[Two-sided preconditioning]
M --> P
P --> U[Parameter update]
Reviews (2): Last reviewed commit: "fix: use decoupled weight decay in OKLS" | Re-trigger Greptile |
https://blog.tilderesearch.com/blog/online-kl-shampoo