ReflectionEraser/ReflectionEraserApp
0
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 