(self, img: torch.Tensor, img_meta: Iterable,
**kwargs)
| 126 | raise NotImplementedError('This method is not implemented.') |
| 127 | |
| 128 | def simple_test(self, img: torch.Tensor, img_meta: Iterable, |
| 129 | **kwargs) -> list: |
| 130 | with torch.cuda.device(self.device_id), torch.no_grad(): |
| 131 | seg_pred = self.model({'input': img})['output'] |
| 132 | seg_pred = seg_pred.detach().cpu().numpy() |
| 133 | # whole might support dynamic reshape |
| 134 | ori_shape = img_meta[0]['ori_shape'] |
| 135 | if not (ori_shape[0] == seg_pred.shape[-2] |
| 136 | and ori_shape[1] == seg_pred.shape[-1]): |
| 137 | seg_pred = torch.from_numpy(seg_pred).float() |
| 138 | seg_pred = resize( |
| 139 | seg_pred, size=tuple(ori_shape[:2]), mode='nearest') |
| 140 | seg_pred = seg_pred.long().detach().cpu().numpy() |
| 141 | seg_pred = seg_pred[0] |
| 142 | seg_pred = list(seg_pred) |
| 143 | return seg_pred |
| 144 | |
| 145 | def aug_test(self, imgs, img_metas, **kwargs): |
| 146 | raise NotImplementedError('This method is not implemented.') |
nothing calls this directly
no outgoing calls
no test coverage detected