| 172 | |
| 173 | |
| 174 | class 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 |
no outgoing calls
no test coverage detected