(self, batch)
| 331 | self.scaler = GradScaler() if cfg.TRAINER.PLOTPP.PREC == "amp" else None |
| 332 | |
| 333 | def forward_backward(self, batch): |
| 334 | image, label = self.parse_batch_train(batch) |
| 335 | |
| 336 | prec = self.cfg.TRAINER.PLOTPP.PREC |
| 337 | if prec == "amp": |
| 338 | with autocast(): |
| 339 | output = self.model(image, label) |
| 340 | self.optim.zero_grad() |
| 341 | self.scaler.scale(output).backward() |
| 342 | self.scaler.step(self.optim) |
| 343 | self.scaler.update() |
| 344 | else: |
| 345 | output = self.model(image) |
| 346 | loss = F.cross_entropy(output, label) |
| 347 | self.model_backward_and_update(loss) |
| 348 | |
| 349 | loss_summary = {"loss": loss.item(), |
| 350 | "acc": compute_accuracy(output, label)[0].item()} |
| 351 | |
| 352 | if (self.batch_idx + 1) == self.num_batches: |
| 353 | self.update_lr() |
| 354 | |
| 355 | return loss_summary |
| 356 | |
| 357 | def parse_batch_train(self, batch): |
| 358 | input = batch["img"] |
nothing calls this directly
no test coverage detected