| 125 | """ |
| 126 | |
| 127 | def __init__(self, params, lr=required, momentum=0, dampening=0, |
| 128 | weight_decay=0, nesterov=False): |
| 129 | if lr is not required and lr < 0.0: |
| 130 | raise ValueError("Invalid learning rate: {}".format(lr)) |
| 131 | if momentum < 0.0: |
| 132 | raise ValueError("Invalid momentum value: {}".format(momentum)) |
| 133 | if weight_decay < 0.0: |
| 134 | raise ValueError( |
| 135 | "Invalid weight_decay value: {}".format(weight_decay)) |
| 136 | |
| 137 | defaults = dict(lr=lr, momentum=momentum, dampening=dampening, |
| 138 | weight_decay=weight_decay, nesterov=nesterov) |
| 139 | if nesterov and (momentum <= 0 or dampening != 0): |
| 140 | raise ValueError( |
| 141 | "Nesterov momentum requires a momentum and zero dampening") |
| 142 | super(SGD, self).__init__(params, defaults) |
| 143 | |
| 144 | def __setstate__(self, state): |
| 145 | super(SGD, self).__setstate__(state) |