MCPcopy Create free account
hub / github.com/awslabs/gap-text2sql / AdamW

Class AdamW

rat-sql-gap/seq2struct/optimizers.py:128–236  ·  view source on GitHub ↗

Implements Adam algorithm. It has been proposed in `Adam: A Method for Stochastic Optimization`_. Arguments: params (iterable): iterable of parameters to optimize or dicts defining parameter groups lr (float, optional): learning rate (default: 1e-3) betas

Source from the content-addressed store, hash-verified

126
127@registry.register('optimizer', 'adamw')
128class AdamW(torch.optim.Optimizer):
129 """Implements Adam algorithm.
130 It has been proposed in `Adam: A Method for Stochastic Optimization`_.
131 Arguments:
132 params (iterable): iterable of parameters to optimize or dicts defining
133 parameter groups
134 lr (float, optional): learning rate (default: 1e-3)
135 betas (Tuple[float, float], optional): coefficients used for computing
136 running averages of gradient and its square (default: (0.9, 0.999))
137 eps (float, optional): term added to the denominator to improve
138 numerical stability (default: 1e-8)
139 weight_decay (float, optional): weight decay (L2 penalty) (default: 0)
140 amsgrad (boolean, optional): whether to use the AMSGrad variant of this
141 algorithm from the paper `On the Convergence of Adam and Beyond`_
142 .. _Adam\: A Method for Stochastic Optimization:
143 https://arxiv.org/abs/1412.6980
144 .. _On the Convergence of Adam and Beyond:
145 https://openreview.net/forum?id=ryQu7f-RZ
146
147 **Modified to implement AdamW, see https://arxiv.org/pdf/1711.05101v3.pdf**
148 """
149
150 def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
151 weight_decay=0, amsgrad=False):
152 if not 0.0 <= lr:
153 raise ValueError("Invalid learning rate: {}".format(lr))
154 if not 0.0 <= eps:
155 raise ValueError("Invalid epsilon value: {}".format(eps))
156 if not 0.0 <= betas[0] < 1.0:
157 raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0]))
158 if not 0.0 <= betas[1] < 1.0:
159 raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1]))
160 defaults = dict(lr=lr, betas=betas, eps=eps,
161 weight_decay=weight_decay, amsgrad=amsgrad)
162 super(AdamW, self).__init__(params, defaults)
163
164 def __setstate__(self, state):
165 super(AdamW, self).__setstate__(state)
166 for group in self.param_groups:
167 group.setdefault('amsgrad', False)
168
169 def step(self, closure=None):
170 """Performs a single optimization step.
171 Arguments:
172 closure (callable, optional): A closure that reevaluates the model
173 and returns the loss.
174 """
175 loss = None
176 if closure is not None:
177 loss = closure()
178
179 for group in self.param_groups:
180 for p in group['params']:
181 if p.grad is None:
182 continue
183 grad = p.grad.data
184 if grad.is_sparse:
185 raise RuntimeError('Adam does not support sparse gradients, please consider SparseAdam instead')

Callers 2

get_optimizersMethod · 0.85
get_optimizersMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected