import torch from torch import Tensor from collections.abc import Iterable classAdamW(torch.optim.Optimizer): def__init__(self,params:Iterable[Tensor]|list[dict], lr: float | Tensor = 1e-3, betas: tuple[float | Tensor, float | Tensor] = (0.9, 0.999), eps: float = 1e-8, weight_decay: float = 0.01): if lr<0.0: raise ValueError("learning rate can't to be a negative number") if betas[0]<0.0or betas[0]>=1.0or betas[1]<0.0or betas[1]>=1.0: raise ValueError("betas need to be in the [0.0,1.0)") if eps<=0.0: raise ValueError("eps need to be a positive number") if weight_decay<0.0: raise ValueError("weight_decay can't to be a negative number") defaults={ "lr":lr, "betas":betas, "eps":eps, "weight_decay":weight_decay } super().__init__(params,defaults)
@torch.no_grad defstep(self,closure=None): loss=None if closure isnotNone: with torch.enable_grad(): loss=closure()
for param_group inself.param_groups: lr=param_group['lr'] betas=param_group['betas'] eps=param_group['eps'] weight_decay=param_group['weight_decay'] for p in param_group['params']: if p.grad isNone: continue state=self.state[p] # 获取状态 iflen(state)==0: step=state.get('step',0) exp_avg=state.get('exp_avg',torch.zeros_like(p,dtype=p.dtype,device=p.device)) exp_avg_sq=state.get('exp_avg_sq',torch.zeros_like(p,dtype=p.dtype,device=p.device)) else: step=state['step'] exp_avg=state['exp_avg'] exp_avg_sq=state['exp_avg_sq'] step+=1 state['step']=step