MCPcopy Create free account
hub / github.com/IGNF/FLAIR-1 / segmentation_task_predict

Class segmentation_task_predict

src/flair/task_module.py:174–213  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

172
173
174class segmentation_task_predict(pl.LightningModule):
175 def __init__(
176 self,
177 model,
178 num_classes : int,
179 use_metadata : bool = False,
180 ):
181
182 super().__init__()
183 self.model = model
184 self.num_classes = num_classes
185 self.use_metadata=use_metadata
186
187 def forward(self, input_im, input_met):
188 logits = self.model(input_im, input_met)
189 return logits
190
191 def step(self, batch):
192 if self.use_metadata == True:
193 images, metadata, targets = batch["img"], batch["mtd"], batch["msk"]
194 else:
195 images, metadata, targets = batch["img"], '', batch["msk"]
196 logits = self.forward(images, metadata)
197
198 with torch.no_grad():
199 proba = torch.softmax(logits, dim=1)
200 preds = torch.argmax(proba, dim=1)
201 targets = torch.argmax(targets, dim=1)
202 preds = preds.flatten(start_dim=1) # Change shapes and cast target to integer for metrics computation
203 targets = targets.flatten(start_dim=1).type(torch.int32)
204 return preds, targets
205
206 def predict_step(self, batch, batch_idx, dataloader_idx=0):
207 if self.use_metadata == True:
208 logits = self.forward(batch["img"], batch["mtd"])
209 else:
210 logits = self.forward(batch["img"], '')
211 proba = torch.softmax(logits, dim=1)
212 batch["preds"] = torch.argmax(proba, dim=1)
213 return batch
214
215

Callers 1

get_segmentation_moduleFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected