Launch validation.
(self)
| 368 | self.val_loss: Dict[str, HistoryBuffer] = dict() |
| 369 | |
| 370 | def run(self) -> dict: |
| 371 | """Launch validation.""" |
| 372 | self.runner.call_hook('before_val') |
| 373 | self.runner.call_hook('before_val_epoch') |
| 374 | self.runner.model.eval() |
| 375 | |
| 376 | # clear val loss |
| 377 | self.val_loss.clear() |
| 378 | for idx, data_batch in enumerate(self.dataloader): |
| 379 | self.run_iter(idx, data_batch) |
| 380 | |
| 381 | # compute metrics |
| 382 | metrics = self.evaluator.evaluate(len(self.dataloader.dataset)) |
| 383 | |
| 384 | if self.val_loss: |
| 385 | loss_dict = _parse_losses(self.val_loss, 'val') |
| 386 | metrics.update(loss_dict) |
| 387 | |
| 388 | self.runner.call_hook('after_val_epoch', metrics=metrics) |
| 389 | self.runner.call_hook('after_val') |
| 390 | return metrics |
| 391 | |
| 392 | @torch.no_grad() |
| 393 | def run_iter(self, idx, data_batch: Sequence[dict]): |
no test coverage detected