Test with augmentations. Only rescale=True is supported.
(self, imgs, img_metas, rescale=True)
| 270 | return seg_logit |
| 271 | |
| 272 | def aug_test(self, imgs, img_metas, rescale=True): |
| 273 | """Test with augmentations. |
| 274 | |
| 275 | Only rescale=True is supported. |
| 276 | """ |
| 277 | # aug_test rescale all imgs back to ori_shape for now |
| 278 | assert rescale |
| 279 | # to save memory, we get augmented seg logit inplace |
| 280 | seg_logit = self.inference(imgs[0], img_metas[0], rescale) |
| 281 | for i in range(1, len(imgs)): |
| 282 | cur_seg_logit = self.inference(imgs[i], img_metas[i], rescale) |
| 283 | seg_logit += cur_seg_logit |
| 284 | seg_logit /= len(imgs) |
| 285 | if self.out_channels == 1: |
| 286 | seg_pred = (seg_logit > self.decode_head.threshold).to(seg_logit).squeeze(1) |
| 287 | else: |
| 288 | seg_pred = seg_logit.argmax(dim=1) |
| 289 | seg_pred = seg_pred.cpu().numpy() |
| 290 | # unravel batch dim |
| 291 | seg_pred = list(seg_pred) |
| 292 | return seg_pred |
| 293 | |
| 294 | def aug_test_logits(self, img, img_metas, rescale=True): |
| 295 | """Test with augmentations. |