From d855c09cf3c881c3c6f89581524daa4cd46cdcea Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sat, 1 Aug 2026 09:20:54 +0800 Subject: [PATCH] fix: use torch.optim.AdamW in ManoAdamW instead of NAdamW - ManoAdamW now uses torch.optim.AdamW(fused=True, betas=(0.9, 0.95)) matching MuonAdamW, eliminating a confounding variable in optimizer comparison experiments - only NoraNAdamW retains NAdamW, which is correct per the Nora paper design --- astrai/optim/mano_adamw.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/astrai/optim/mano_adamw.py b/astrai/optim/mano_adamw.py index ba6cc3a..bc2dd89 100644 --- a/astrai/optim/mano_adamw.py +++ b/astrai/optim/mano_adamw.py @@ -10,7 +10,7 @@ Reference: https://arxiv.org/abs/2601.23000 import math import torch -from torch import nn +from torch import nn, optim from torch.optim import Optimizer from astrai.optim.composite import ( @@ -20,7 +20,7 @@ from astrai.optim.composite import ( composite_zero_grad, refresh_param_groups, ) -from astrai.optim.nora_nadamw import NAdamW, partition_optimizer_parameters +from astrai.optim.nora_nadamw import partition_optimizer_parameters class Mano(Optimizer): @@ -162,7 +162,16 @@ class ManoAdamW(Optimizer): ) if groups.nadamw_no_decay: adamw_groups.append({"params": groups.nadamw_no_decay, "weight_decay": 0.0}) - self.adamw = NAdamW(adamw_groups, lr=lr) if adamw_groups else None + self.adamw = ( + optim.AdamW( + adamw_groups, + lr=lr, + betas=(0.9, 0.95), + fused=True, + ) + if adamw_groups + else None + ) self.param_groups = refresh_param_groups([self.mano, self.adamw]) @torch.no_grad()