MCPcopy Create free account
hub / github.com/CompVis/diff2flow / DepthMetricTracker

Class DepthMetricTracker

diff2flow/metrics.py:62–101  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

60
61
62class DepthMetricTracker:
63 def __init__(self):
64 super().__init__()
65 self.delta1s = []
66 self.relabs = []
67
68 self.quantiles = (0.795, 46.679) # DIODE quantiles
69
70 def __call__(self, target, pred):
71 """ Assumes target and pred in [-1, 1] range """
72 target = target.mean(dim=1, keepdim=True)
73 pred = pred.mean(dim=1, keepdim=True)
74
75 # unnormalize ground truth depth
76 target = target / 2 + 0.5
77 q_min, q_max = self.quantiles
78 q_min = torch.log(torch.tensor(q_min))
79 q_max = torch.log(torch.tensor(q_max))
80 log_target = target * (q_max - q_min) + q_min
81 target = torch.exp(log_target)
82
83 # scale and shift invariance
84 pred = apply_scale_and_shift(pred=pred, gt=log_target)
85 pred = pred.exp()
86
87 # compute metrics
88 relabs = abs_rel_error(pred, target)
89 self.relabs.append(relabs)
90 delta1 = delta1_accuracy(pred, target)
91 self.delta1s.append(delta1)
92
93 def reset(self):
94 self.delta1s = []
95 self.relabs = []
96
97 def aggregate(self):
98 delta1 = torch.stack(self.delta1s).mean()
99 relabs = torch.stack(self.relabs).mean()
100 out = dict(delta1=delta1, relabs=relabs)
101 return out

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected