U Ph-@sXddlmZddlmZmZddlmZddlZddlm Z edZ Gddde Z dS) ) annotations)CallableIterable)TypeVarN) OptimizerTc sReZdZdZdddd ddd d d fd d ZfddZddddddZZS)Novograda Novograd based on `Stochastic Gradient Methods with Layer-wise Adaptive Moments for Training of Deep Networks `_. The code is adapted from the implementations in `Jasper for PyTorch `_, and `OpenSeq2Seq `_. Args: params: iterable of parameters to optimize or dicts defining parameter groups. lr: learning rate. Defaults to 1e-3. betas: coefficients used for computing running averages of gradient and its square. Defaults to (0.9, 0.98). eps: term added to the denominator to improve numerical stability. Defaults to 1e-8. weight_decay: weight decay (L2 penalty). Defaults to 0. grad_averaging: gradient averaging. Defaults to ``False``. amsgrad: whether to use the AMSGrad variant of this algorithm from the paper `On the Convergence of Adam and Beyond `_. Defaults to ``False``. MbP?g?g\(\?:0yE>rFrfloatztuple[float, float]bool)paramslrbetaseps weight_decaygrad_averagingamsgradc sd|krtd|d|kr,td|d|dkrDdksXntd|dd|dkrpdksntd|dd|krtd |t||||||d }t||dS) NgzInvalid learning rate: zInvalid epsilon value: rg?z#Invalid beta parameter at index 0: z#Invalid beta parameter at index 1: zInvalid weight_decay value: )rrrrrr) ValueErrordictsuper__init__) selfrrrrrrrdefaults __class__N/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/optimizers/novograd.pyr*s& zNovograd.__init__cs(t||jD]}|ddqdS)NrF)r __setstate__ param_groups setdefault)rstategrouprrrr Ds  zNovograd.__setstate__NzCallable[[], T] | NonezT | None)closurereturncCsd}|dk r|}|jD]}|dD]}|jdkr8q&|jj}|jrNtd|d}|j|}t|dkrd|d<t|j|d<t g |dj |d<|rt g |dj |d <|d|d}} |r|d } |d \} } |dd 7<t t |d } | dkr| | n| | j| d | d |r`tj| | | d| |d}n| |d}|||ddkr|j|j|dd |dr|d | || ||jj||d d q&q|S)zPerforms a single optimization step. Arguments: closure: A closure that reevaluates the model and returns the loss. Defaults to ``None``. Nrz#Sparse gradients are not supported.rrstepexp_avg exp_avg_sqmax_exp_avg_sqrr)alpha)outrrrr)r!graddata is_sparse RuntimeErrorr#lentorch zeros_likezerostodevicesumpowcopy_mul_add_maxsqrtdiv_)rr%lossr$pr.rr#r(r)r*beta1beta2normdenomrrrr'IsN         z Novograd.step)r r r rFF)N)__name__ __module__ __qualname____doc__rr r' __classcell__rrrrrs  r) __future__rcollections.abcrrtypingrr3 torch.optimrrrrrrr s