| 5 | |
| 6 | |
| 7 | class SitsScdModel(L.LightningModule): |
| 8 | def __init__(self, cfg): |
| 9 | super().__init__() |
| 10 | self.cfg = cfg |
| 11 | self.model = instantiate(cfg.network.instance) |
| 12 | self.loss = instantiate(cfg.loss.instance) |
| 13 | self.ignore_index = self.loss.ignore_index |
| 14 | self.val_metrics = {'out': instantiate(cfg.val_metrics), 'in': instantiate(cfg.val_metrics)} |
| 15 | self.test_metrics = {'out': instantiate(cfg.test_metrics), 'in': instantiate(cfg.test_metrics)} |
| 16 | self.domain_dict = {0: 'out', 1: 'in'} |
| 17 | |
| 18 | def training_step(self, batch, batch_idx): |
| 19 | pred = self.model(batch) |
| 20 | loss = self.loss(pred, batch, average=True) |
| 21 | for metric_name, metric_value in loss.items(): |
| 22 | self.log( |
| 23 | f"train/{metric_name}", |
| 24 | metric_value, |
| 25 | sync_dist=True, |
| 26 | on_step=True, |
| 27 | on_epoch=True, |
| 28 | ) |
| 29 | return loss |
| 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']: |
| 41 | metrics = self.val_metrics[dataloader_idx].compute() |
| 42 | for metric_name, metric_value in metrics.items(): |
| 43 | self.log( |
| 44 | f"val/{metric_name}_{dataloader_idx}", |
| 45 | metric_value, |
| 46 | sync_dist=True, |
| 47 | on_step=False, |
| 48 | on_epoch=True, |
| 49 | ) |
| 50 | |
| 51 | @torch.no_grad() |
| 52 | def test_step(self, batch, batch_idx, dataloader_idx): |
| 53 | pred = self.model(batch) |
| 54 | pred["pred"] = torch.argmax(pred["logits"], dim=2) |
| 55 | self.test_metrics[self.domain_dict[dataloader_idx]].update(pred["pred"], batch["gt"]) |
| 56 | |
| 57 | def on_test_epoch_end(self): |
| 58 | for dataloader_idx in ['out', 'in']: |
| 59 | metrics = self.test_metrics[dataloader_idx].compute() |
| 60 | for metric_name, metric_value in metrics.items(): |
| 61 | self.log( |
| 62 | f"test/{metric_name}_{dataloader_idx}", |
| 63 | metric_value, |
| 64 | sync_dist=True, |