o / i-@sXddlmZddlmZmZddlmZddlZddlm Z edZ Gddde Z dS) ) annotations)CallableIterable)TypeVarN) OptimizerTcsHeZdZdZ      ddfdd ZfddZdd ddZZS)!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>rFparamsrlrfloatbetastuple[float, float]eps weight_decaygrad_averagingboolamsgradc sd|kr td|d|krtd|d|dkr"dks,ntd|dd|dkr8dksBntd|dd|krMtd |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: )r rrrrr) ValueErrordictsuper__init__) selfr r rrrrrdefaults __class__[/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/optimizers/novograd.pyr*s  zNovograd.__init__cs(t||jD]}|ddq dS)NrF)r __setstate__ param_groups setdefault)rstategrouprrr r!Ds  zNovograd.__setstate__NclosureCallable[[], T] | NonereturnT | NonecCsd}|dur |}|jD]}|dD]}|jdurq|jj}|jr%td|d}|j|}t|dkr\d|d<t|j|d<t g |dj |d<|r\t g |dj |d <|d|d}} |rk|d } |d \} } |dd 7<t t |d } | dkr| | n | | j| d | d |rtj| | | d| |d}n | |d}|||ddkr|j|j|dd |dr|d | || ||jj||d d qq |S)zPerforms a single optimization step. Arguments: closure: A closure that reevaluates the model and returns the loss. Defaults to ``None``. Nr z#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%pr1rr$r+r,r-beta1beta2normdenomrrr r*IsP         4z Novograd.step)r r r rFF)r rr rrrrrrrrrrr)N)r&r'r(r))__name__ __module__ __qualname____doc__rr!r* __classcell__rrrr rs r) __future__rcollections.abcrrtypingrr6 torch.optimrrrrrrr s