(self, batch, batch_idx, dataloader_idx)
| 30 | |
| 31 | @torch.no_grad() |
| 32 | def validation_step(self, batch, batch_idx, dataloader_idx): |
| 33 | pred = self.model(batch) |
| 34 | pred["pred"] = torch.argmax(pred["logits"], dim=2) |
| 35 | loss = self.loss(pred, batch, average=True)["loss"] |
| 36 | self.val_metrics[self.domain_dict[dataloader_idx]].update(pred["pred"], batch["gt"]) |
| 37 | self.log("val/loss", loss, sync_dist=True, on_step=False, on_epoch=True) |
| 38 | |
| 39 | def on_validation_epoch_end(self): |
| 40 | for dataloader_idx in ['out', 'in']: |