Move optimizer to a specific device. Args: optimizer: the optimzer to move device: the torch device
(optimizer: torch.optim.Optimizer, device: torch.device)
| 6 | |
| 7 | |
| 8 | def optimizer_to(optimizer: torch.optim.Optimizer, device: torch.device): |
| 9 | """ |
| 10 | Move optimizer to a specific device. |
| 11 | |
| 12 | Args: |
| 13 | optimizer: |
| 14 | the optimzer to move |
| 15 | device: |
| 16 | the torch device |
| 17 | """ |
| 18 | for param in optimizer.state.values(): |
| 19 | if isinstance(param, torch.Tensor): |
| 20 | param.data = param.data.to(device) |
| 21 | if param._grad is not None: |
| 22 | param._grad.data = param._grad.data.to(device) |
| 23 | elif isinstance(param, dict): |
| 24 | for subparam in param.values(): |
| 25 | if isinstance(subparam, torch.Tensor): |
| 26 | subparam.data = subparam.data.to(device) |
| 27 | if subparam._grad is not None: |
| 28 | subparam._grad.data = subparam._grad.data.to(device) |