MCPcopy Create free account
hub / github.com/ElliotVincent/SitsSCD / SitsScdModel

Class SitsScdModel

models/module.py:7–102  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5
6
7class 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,

Callers 1

load_modelFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected