Args: model, data_loader, optimizer: same as in :class:`SimpleTrainer`. grad_scaler: torch GradScaler to automatically scale gradients.
(self, model, data_loader, optimizer, grad_scaler=None)
| 293 | """ |
| 294 | |
| 295 | def __init__(self, model, data_loader, optimizer, grad_scaler=None): |
| 296 | """ |
| 297 | Args: |
| 298 | model, data_loader, optimizer: same as in :class:`SimpleTrainer`. |
| 299 | grad_scaler: torch GradScaler to automatically scale gradients. |
| 300 | """ |
| 301 | unsupported = "AMPTrainer does not support single-process multi-device training!" |
| 302 | if isinstance(model, DistributedDataParallel): |
| 303 | assert not (model.device_ids and len(model.device_ids) > 1), unsupported |
| 304 | assert not isinstance(model, DataParallel), unsupported |
| 305 | |
| 306 | super().__init__(model, data_loader, optimizer) |
| 307 | |
| 308 | if grad_scaler is None: |
| 309 | from torch.cuda.amp import GradScaler |
| 310 | |
| 311 | grad_scaler = GradScaler() |
| 312 | self.grad_scaler = grad_scaler |
| 313 | |
| 314 | def run_step(self): |
| 315 | """ |