-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathutils.py
More file actions
184 lines (144 loc) · 6.48 KB
/
Copy pathutils.py
File metadata and controls
184 lines (144 loc) · 6.48 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
"""Evaluation metrics and lightweight logging/reporting utilities."""
import os
import sys
from pathlib import Path
import numpy as np
import torch
import lpips
from pytorch_fid import fid_score
from torchmetrics.image import (
PeakSignalNoiseRatio,
StructuralSimilarityIndexMeasure,
SpatialCorrelationCoefficient,
)
from torchmetrics.image.inception import InceptionScore
def write(log, string):
"""Flush stdout and append a line to a log file handle."""
sys.stdout.flush()
log.write(string + '\n')
log.flush()
class ImageQualityMetrics:
"""Wraps the image-quality metrics reported in the paper.
Note:
``denormalize`` maps model outputs to [0, 1] with ``img + 0.5``. This is
the exact transform used to produce the reported numbers; change it only
if you also re-generate all baselines for a fair comparison.
"""
def __init__(self, device='cuda', fid_real_images_path=None):
self.device = device
self.psnr = PeakSignalNoiseRatio(data_range=1.0).to(device)
self.ssim = StructuralSimilarityIndexMeasure(data_range=1.0).to(device)
self.scc = SpatialCorrelationCoefficient().to(device)
self.lpips_fn = lpips.LPIPS(net='vgg').to(device)
self.inception_score_fn = InceptionScore().to(device)
# FID requires a directory of real images.
self.fid_real_images_path = fid_real_images_path
def denormalize(self, img):
"""Map an image to the [0, 1] range used for metric computation."""
return (img + 0.5).clamp(0, 1)
def denormalize_and_convert(self, img):
"""Denormalize and convert to uint8 (required by Inception Score)."""
img = self.denormalize(img)
return (img * 255).to(torch.uint8)
def calculate_psnr(self, real_images, generated_images):
return self.psnr(self.denormalize(generated_images), self.denormalize(real_images))
def calculate_ssim(self, real_images, generated_images):
return self.ssim(self.denormalize(generated_images), self.denormalize(real_images))
def calculate_scc(self, real_images, generated_images):
return self.scc(self.denormalize(generated_images), self.denormalize(real_images))
def calculate_lpips(self, real_images, generated_images):
return self.lpips_fn(
self.denormalize(generated_images), self.denormalize(real_images)
).mean()
def calculate_fid(self, generated_images_path):
if not self.fid_real_images_path:
raise ValueError("Path to real images for FID calculation not provided.")
return fid_score.calculate_fid_given_paths(
[self.fid_real_images_path, generated_images_path],
batch_size=50, device=self.device, dims=2048,
)
def calculate_inception_score(self, generated_images):
gen_images_uint8 = self.denormalize_and_convert(generated_images)
return self.inception_score_fn(gen_images_uint8)
def calculate_metrics(self, real_images, generated_images, generated_images_path=None):
"""Compute PSNR/SSIM/SCC/LPIPS (+ optional IS and FID) for a batch."""
metrics = {
'psnr': self.calculate_psnr(real_images, generated_images).item(),
'ssim': self.calculate_ssim(real_images, generated_images).item(),
'scc': self.calculate_scc(real_images, generated_images).item(),
'lpips': self.calculate_lpips(real_images, generated_images).item(),
}
# Inception Score is only meaningful for a sufficiently large batch.
if generated_images.shape[0] > 10:
is_mean, is_std = self.calculate_inception_score(generated_images)
metrics['is_mean'] = is_mean.item()
metrics['is_std'] = is_std.item()
else:
metrics['is_mean'] = 0
metrics['is_std'] = 0
if generated_images_path:
metrics['fid'] = self.calculate_fid(generated_images_path).item()
else:
metrics['fid'] = 0
return metrics
class Report:
"""Append-mode text logger that also echoes to stdout."""
def __init__(self, save_dir, type):
filename = os.path.join(save_dir, f'{type}_log.txt')
if not os.path.exists(save_dir):
Path(save_dir).mkdir(parents=True, exist_ok=True)
mode = 'a' if os.path.exists(filename) else 'w'
self.logFile = open(filename, mode)
def write(self, string):
print(string)
write(self.logFile, string)
def __del__(self):
self.logFile.close()
class Train_Report:
"""Accumulates the running training loss over a logging interval."""
def __init__(self):
self.total_loss = []
self.num_examples = 0
def update(self, batch_size, total_loss):
self.num_examples += batch_size
self.total_loss.append(total_loss * batch_size)
def compute_mean(self):
self.total_loss = np.sum(self.total_loss) / self.num_examples
def result_str(self, lr, period_time):
self.compute_mean()
return (f'Total Loss: {self.total_loss:.6f}\t'
f'learning rate: {lr:.7f}\tTime: {period_time:.4f}')
class Test_Report:
"""Accumulates per-batch metrics and reports their means."""
def __init__(self):
self.psnr = []
self.ssim = []
self.scc = []
self.lpips = []
self.fid = []
self.is_mean = []
self.is_std = []
self.num_examples = 0
def update(self, batch_size, metrics):
self.num_examples += batch_size
self.psnr.append(metrics['psnr'])
self.ssim.append(metrics['ssim'])
self.scc.append(metrics['scc'])
self.lpips.append(metrics['lpips'])
self.fid.append(metrics['fid'])
self.is_mean.append(metrics['is_mean'])
self.is_std.append(metrics['is_std'])
def compute_mean(self):
self.psnr = np.sum(self.psnr) / self.num_examples
self.ssim = np.sum(self.ssim) / self.num_examples
self.scc = np.sum(self.scc) / self.num_examples
self.lpips = np.sum(self.lpips) / self.num_examples
self.fid = np.sum(self.fid) / self.num_examples
self.is_mean = np.sum(self.is_mean) / self.num_examples
self.is_std = np.sum(self.is_std) / self.num_examples
def result_str(self):
self.compute_mean()
return (f'PSNR: {self.psnr:.6f}\tSSIM: {self.ssim:.6f}\t'
f'SCC: {self.scc:.6f}\tLPIPS: {self.lpips:.6f}\t'
f'FID: {self.fid:.6f}\tIS-mean: {self.is_mean:.6f}\t'
f'IS-std: {self.is_std:.6f}')