• torch.optim.Adam() 函数用法


    Adam: A method for stochastic optimization

     Adam是通过梯度的一阶矩和二阶矩自适应的控制每个参数的学习率的大小。

     adam的初始化

    1. def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
    2. weight_decay=0, amsgrad=False):
    Args:
        params (iterable): iterable of parameters to optimize or dicts defining
            parameter groups
        lr (float, optional): learning rate (default: 1e-3)
        betas (Tuple[float, float], optional): coefficients used for computing
            running averages of gradient and its square (default: (0.9, 0.999))
        eps (float, optional): term added to the denominator to improve
            numerical stability (default: 1e-8)
        weight_decay (float, optional): weight decay (L2 penalty) (default: 0)
        amsgrad (boolean, optional): whether to use the AMSGrad variant of this
            algorithm from the paper `On the Convergence of Adam and Beyond`_
            (default: False)
    1. @torch.no_grad()
    2. def step(self, closure=None):
    3. """Performs a single optimization step.
    4. Args:
    5. closure (callable, optional): A closure that reevaluates the model
    6. and returns the loss.
    7. """
    8. loss = None
    9. if closure is not None:
    10. with torch.enable_grad():
    11. loss = closure()
    12. for group in self.param_groups:
    13. params_with_grad = []
    14. grads = []
    15. exp_avgs = []
    16. exp_avg_sqs = []
    17. max_exp_avg_sqs = []
    18. state_steps = []
    19. beta1, beta2 = group['betas']
    20. for p in group['params']:
    21. if p.grad is not None:
    22. params_with_grad.append(p)
    23. if p.grad.is_sparse:
    24. raise RuntimeError('Adam does not support sparse gradients, please consider SparseAdam instead')
    25. grads.append(p.grad)
    26. state = self.state[p]
    27. # Lazy state initialization
    28. if len(state) == 0:
    29. state['step'] = 0
    30. # Exponential moving average of gradient values
    31. state['exp_avg'] = torch.zeros_like(p, memory_format=torch.preserve_format)
    32. # Exponential moving average of squared gradient values
    33. state['exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format)
    34. if group['amsgrad']:
    35. # Maintains max of all exp. moving avg. of sq. grad. values
    36. state['max_exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format)
    37. exp_avgs.append(state['exp_avg'])
    38. exp_avg_sqs.append(state['exp_avg_sq'])
    39. if group['amsgrad']:
    40. max_exp_avg_sqs.append(state['max_exp_avg_sq'])
    41. # update the steps for each param group update
    42. state['step'] += 1
    43. # record the step after step update
    44. state_steps.append(state['step'])
    45. F.adam(params_with_grad,
    46. grads,
    47. exp_avgs,
    48. exp_avg_sqs,
    49. max_exp_avg_sqs,
    50. state_steps,
    51. amsgrad=group['amsgrad'],
    52. beta1=beta1,
    53. beta2=beta2,
    54. lr=group['lr'],
    55. weight_decay=group['weight_decay'],
    56. eps=group['eps'])
    57. return loss

     

  • 相关阅读:
    科研TCO-PEG-alginate|反式环辛烯-聚乙二醇-海藻酸钠|TCO-PEG-海藻酸钠
    15 Go的并发
    win10换ubuntu
    操作系统的发展与分类
    Java异常处理机制
    C++中指针指向无效的内存单元
    分享一个实用的MySQL一键巡检脚本
    Java Socket实现简易多人聊天室传输聊天内容或文件
    Optional 详解
    elasticsearch
  • 原文地址:https://blog.csdn.net/qq_40107571/article/details/126018026