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

Class ImageMetricTracker

diff2flow/metrics.py:20–59  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

18
19
20class ImageMetricTracker(nn.Module):
21 def __init__(self):
22 super().__init__()
23
24 self.ssim = SSIM(data_range=1.)
25 self.ssims = []
26
27 self.psnrs = []
28
29 self.fid = FrechetInceptionDistance(
30 feature=2048,
31 reset_real_features=True,
32 normalize=False,
33 sync_on_compute=True
34 )
35
36 def __call__(self, target, pred):
37 """ Assumes target and pred in [-1, 1] range """
38 real_ims = un_normalize_ims(target)
39 fake_ims = un_normalize_ims(pred)
40
41 # update FID
42 self.fid.update(real_ims, real=True)
43 self.fid.update(fake_ims, real=False)
44
45 # SSIM and PSNR
46 self.ssims.append(self.ssim(pred/2+0.5, target/2+0.5))
47 self.psnrs.append(calculate_PSNR(pred/2+0.5, target/2+0.5))
48
49 def reset(self):
50 self.ssims = []
51 self.psnrs = []
52 self.fid.reset()
53
54 def aggregate(self):
55 fid = self.fid.compute()
56 ssim = torch.stack(self.ssims).mean()
57 psnr = torch.stack(self.psnrs).mean()
58 out = dict(fid=fid, ssim=ssim, psnr=psnr)
59 return out
60
61
62class DepthMetricTracker:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected