ReflectionEraser/ReflectionEraserApp
0
1# Metrics/Indexes
2from skimage.metrics import peak_signal_noise_ratio as compare_psnr
3from skimage.metrics import structural_similarity as compare_ssim
4from functools import partial
5import numpy as np
6
7
8class Bandwise(object):
9 def __init__(self, index_fn):
10 self.index_fn = index_fn
11
12 def __call__(self, X, Y):
13 C = X.shape[-1]
14 bwindex = []
15 for ch in range(C):
16 x = X[..., ch]
17 y = Y[..., ch]
18 index = self.index_fn(x, y)
19 bwindex.append(index)
20 return bwindex
21
22
23cal_bwpsnr = Bandwise(partial(compare_psnr, data_range=255))
24cal_bwssim = Bandwise(partial(compare_ssim, data_range=255))
25
26
27def compare_ncc(x, y):
28 return np.mean((x - np.mean(x)) * (y - np.mean(y))) / (np.std(x) * np.std(y))
29
30
31def ssq_error(correct, estimate):
32 """Compute the sum-squared-error for an image, where the estimate is
33 multiplied by a scalar which minimizes the error. Sums over all pixels
34 where mask is True. If the inputs are color, each color channel can be
35 rescaled independently."""
36 assert correct.ndim == 2
37 if np.sum(estimate ** 2) > 1e-5:
38 alpha = np.sum(correct * estimate) / np.sum(estimate ** 2)
39 else:
40 alpha = 0.
41 return np.sum((correct - alpha * estimate) ** 2)
42
43
44def local_error(correct, estimate, window_size, window_shift):
45 """Returns the sum of the local sum-squared-errors, where the estimate may
46 be rescaled within each local region to minimize the error. The windows are
47 window_size x window_size, and they are spaced by window_shift."""
48 M, N, C = correct.shape
49 ssq = total = 0.
50 for c in range(C):
51 for i in range(0, M - window_size + 1, window_shift):
52 for j in range(0, N - window_size + 1, window_shift):
53 correct_curr = correct[i:i + window_size, j:j + window_size, c]
54 estimate_curr = estimate[i:i + window_size, j:j + window_size, c]
55 ssq += ssq_error(correct_curr, estimate_curr)
56 total += np.sum(correct_curr ** 2)
57 # assert np.isnan(ssq/total)
58 return ssq / total
59
60
61def quality_assess(X, Y):
62 # Y: correct; X: estimate
63 psnr = np.mean(cal_bwpsnr(Y, X))
64 ssim = np.mean(cal_bwssim(Y, X))
65 lmse = local_error(Y, X, 20, 10)
66 ncc = compare_ncc(Y, X)
67 return {'PSNR': psnr, 'SSIM': ssim, 'LMSE': lmse, 'NCC': ncc}
68 