Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
losses.py137 linesDownload Raw Back to models
1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4
5from models.vgg import Vgg19
6
7
8def compute_grad(img):
9    gradx = img[..., 1:, :] - img[..., :-1, :]
10    grady = img[..., 1:] - img[..., :-1]
11    return gradx, grady
12
13
14class GradientLoss(nn.Module):
15    def __init__(self):
16        super(GradientLoss, self).__init__()
17        self.loss = nn.L1Loss()
18
19    def forward(self, predict, target):
20        predict_gradx, predict_grady = compute_grad(predict)
21        target_gradx, target_grady = compute_grad(target)
22
23        return self.loss(predict_gradx, target_gradx) + self.loss(predict_grady, target_grady)
24
25
26class MultipleLoss(nn.Module):
27    def __init__(self, losses, weight=None):
28        super(MultipleLoss, self).__init__()
29        self.losses = nn.ModuleList(losses)
30        self.weight = weight or [1 / len(self.losses)] * len(self.losses)
31
32    def forward(self, predict, target):
33        total_loss = 0
34        for weight, loss in zip(self.weight, self.losses):
35            total_loss += loss(predict, target) * weight
36        return total_loss
37
38
39class MeanShift(nn.Conv2d):
40    def __init__(self, data_mean, data_std, data_range=1, norm=True):
41        """norm (bool): normalize/denormalize the stats"""
42        c = len(data_mean)
43        super(MeanShift, self).__init__(c, c, kernel_size=1)
44        std = torch.Tensor(data_std)
45        self.weight.data = torch.eye(c).view(c, c, 1, 1)
46        if norm:
47            self.weight.data.div_(std.view(c, 1, 1, 1))
48            self.bias.data = -1 * data_range * torch.Tensor(data_mean)
49            self.bias.data.div_(std)
50        else:
51            self.weight.data.mul_(std.view(c, 1, 1, 1))
52            self.bias.data = data_range * torch.Tensor(data_mean)
53        self.requires_grad = False
54
55
56class VGGLoss(nn.Module):
57    def __init__(self, vgg=None, weights=None, indices=None, normalize=True):
58        super(VGGLoss, self).__init__()
59        if vgg is None:
60            self.vgg = Vgg19().cuda()
61        else:
62            self.vgg = vgg
63        self.criterion = nn.L1Loss()
64        self.weights = weights or [1.0 / 2.6, 1.0 / 4.8, 1.0 / 3.7, 1.0 / 5.6, 10 / 1.5]
65        self.indices = indices or [2, 7, 12, 21, 30]
66        if normalize:
67            self.normalize = MeanShift([0.485, 0.456, 0.406], [0.229, 0.224, 0.225], norm=True).cuda()
68        else:
69            self.normalize = None
70
71    def forward(self, x, y):
72        if self.normalize is not None:
73            x = self.normalize(x)
74            y = self.normalize(y)
75        x_vgg, y_vgg = self.vgg(x, self.indices), self.vgg(y, self.indices)
76        loss = 0
77        for i in range(len(x_vgg)):
78            loss += self.weights[i] * self.criterion(x_vgg[i], y_vgg[i].detach())
79
80        return loss
81
82
83class ReconsLoss(nn.Module):
84    def __init__(self):
85        super().__init__()
86        self.criterion = nn.L1Loss()
87
88    def forward(self, out_t, out_r, out_rr, input_i):
89        content_diff = self.criterion(out_t + out_r + out_rr, input_i)
90        return content_diff
91
92
93class ExclusionLoss(nn.Module):
94    def __init__(self, level=3, eps=1e-6):
95        super().__init__()
96        self.level = level
97        self.eps = eps
98
99    def forward(self, img_T, img_R):
100        grad_x_loss = []
101        grad_y_loss = []
102
103        for l in range(self.level):
104            grad_x_T, grad_y_T = compute_grad(img_T)
105            grad_x_R, grad_y_R = compute_grad(img_R)
106
107            alphax = (2.0 * torch.mean(torch.abs(grad_x_T))) / (torch.mean(torch.abs(grad_x_R)) + self.eps)
108            alphay = (2.0 * torch.mean(torch.abs(grad_y_T))) / (torch.mean(torch.abs(grad_y_R)) + self.eps)
109
110            gradx1_s = (torch.sigmoid(grad_x_T) * 2) - 1  # mul 2 minus 1 is to change sigmoid into tanh
111            grady1_s = (torch.sigmoid(grad_y_T) * 2) - 1
112            gradx2_s = (torch.sigmoid(grad_x_R * alphax) * 2) - 1
113            grady2_s = (torch.sigmoid(grad_y_R * alphay) * 2) - 1
114
115            grad_x_loss.append(((torch.mean(torch.mul(gradx1_s.pow(2), gradx2_s.pow(2)))) + self.eps) ** 0.25)
116            grad_y_loss.append(((torch.mean(torch.mul(grady1_s.pow(2), grady2_s.pow(2)))) + self.eps) ** 0.25)
117
118            img_T = F.interpolate(img_T, scale_factor=0.5, mode='bilinear')
119            img_R = F.interpolate(img_R, scale_factor=0.5, mode='bilinear')
120        loss_gradxy = torch.sum(sum(grad_x_loss) / 3) + torch.sum(sum(grad_y_loss) / 3)
121
122        return loss_gradxy / 2
123
124
125def init_loss(opt):
126    loss_dic = {}
127    pixel_loss = MultipleLoss([nn.MSELoss(), GradientLoss()], [0.3, 0.6])
128    loss_dic['t_pixel'] = pixel_loss
129    loss_dic['r_pixel'] = pixel_loss
130    loss_dic['recons'] = ReconsLoss()
131    loss_dic['exclu'] = ExclusionLoss(level=3)
132    return loss_dic
133
134
135if __name__ == '__main__':
136    x = torch.randn(3, 32, 224, 224).cuda()
137