Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
index.py68 linesDownload Raw Back to util
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