(pred, target, valid_mask=None)
| 82 | |
| 83 | |
| 84 | def abs_rel_error(pred, target, valid_mask=None): |
| 85 | if pred.shape[1] == 3: |
| 86 | pred = pred.mean(dim=1, keepdim=True) |
| 87 | if target.shape[1] == 3: |
| 88 | target = target.mean(dim=1, keepdim=True) |
| 89 | if valid_mask is not None: |
| 90 | pred = pred[valid_mask] |
| 91 | target = target[valid_mask] |
| 92 | return torch.mean(torch.abs(pred - target) / target) |
| 93 | |
| 94 | |
| 95 | def delta1_accuracy(pred, target): |