Updates the internal evaluation result. Parameters ---------- labels : 'NumpyArray' or list of `NumpyArray` The labels of the data. preds : 'NumpyArray' or list of `NumpyArray` Predicted values.
(self, preds, labels)
| 17 | self.reset() |
| 18 | |
| 19 | def update(self, preds, labels): |
| 20 | """Updates the internal evaluation result. |
| 21 | |
| 22 | Parameters |
| 23 | ---------- |
| 24 | labels : 'NumpyArray' or list of `NumpyArray` |
| 25 | The labels of the data. |
| 26 | preds : 'NumpyArray' or list of `NumpyArray` |
| 27 | Predicted values. |
| 28 | """ |
| 29 | |
| 30 | def evaluate_worker(self, pred, label): |
| 31 | correct, labeled = batch_pix_accuracy(pred, label) |
| 32 | inter, union = batch_intersection_union(pred, label, self.nclass) |
| 33 | |
| 34 | self.total_correct += correct |
| 35 | self.total_label += labeled |
| 36 | if self.total_inter.device != inter.device: |
| 37 | self.total_inter = self.total_inter.to(inter.device) |
| 38 | self.total_union = self.total_union.to(union.device) |
| 39 | self.total_inter += inter |
| 40 | self.total_union += union |
| 41 | |
| 42 | if isinstance(preds, torch.Tensor): |
| 43 | evaluate_worker(self, preds, labels) |
| 44 | elif isinstance(preds, (list, tuple)): |
| 45 | for (pred, label) in zip(preds, labels): |
| 46 | evaluate_worker(self, pred, label) |
| 47 | |
| 48 | def get(self): |
| 49 | """Gets the current evaluation result. |
no outgoing calls
no test coverage detected