| 237 | """ |
| 238 | |
| 239 | def __init__(self, params, lr=required, momentum=0, dampening=0, |
| 240 | weight_decay=0, nesterov=False): |
| 241 | if lr is not required and lr < 0.0: |
| 242 | raise ValueError("Invalid learning rate: {}".format(lr)) |
| 243 | if momentum < 0.0: |
| 244 | raise ValueError("Invalid momentum value: {}".format(momentum)) |
| 245 | if weight_decay < 0.0: |
| 246 | raise ValueError( |
| 247 | "Invalid weight_decay value: {}".format(weight_decay)) |
| 248 | |
| 249 | defaults = dict(lr=lr, momentum=momentum, dampening=dampening, |
| 250 | weight_decay=weight_decay, nesterov=nesterov) |
| 251 | if nesterov and (momentum <= 0 or dampening != 0): |
| 252 | raise ValueError( |
| 253 | "Nesterov momentum requires a momentum and zero dampening") |
| 254 | super(SGDVec, self).__init__(params, defaults) |
| 255 | |
| 256 | def __setstate__(self, state): |
| 257 | super(SGDVec, self).__setstate__(state) |