Scrappy-Doo commited on
Commit
dcda68e
·
verified ·
1 Parent(s): 8fe4632

Upload 2 files

Browse files
basicsr/ops/upfirdn2d/__init__.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ from .upfirdn2d import upfirdn2d
2
+
3
+ __all__ = ['upfirdn2d']
basicsr/ops/upfirdn2d/upfirdn2d.py ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # modify from https://github.com/rosinality/stylegan2-pytorch/blob/master/op/upfirdn2d.py # noqa:E501
2
+
3
+ import torch
4
+ from torch.autograd import Function
5
+ from torch.nn import functional as F
6
+
7
+ try:
8
+ from . import upfirdn2d_ext
9
+ except ImportError:
10
+ import os
11
+ BASICSR_JIT = os.getenv('BASICSR_JIT')
12
+ if BASICSR_JIT == 'True':
13
+ from torch.utils.cpp_extension import load
14
+ module_path = os.path.dirname(__file__)
15
+ upfirdn2d_ext = load(
16
+ 'upfirdn2d',
17
+ sources=[
18
+ os.path.join(module_path, 'src', 'upfirdn2d.cpp'),
19
+ os.path.join(module_path, 'src', 'upfirdn2d_kernel.cu'),
20
+ ],
21
+ )
22
+
23
+
24
+ class UpFirDn2dBackward(Function):
25
+
26
+ @staticmethod
27
+ def forward(ctx, grad_output, kernel, grad_kernel, up, down, pad, g_pad, in_size, out_size):
28
+
29
+ up_x, up_y = up
30
+ down_x, down_y = down
31
+ g_pad_x0, g_pad_x1, g_pad_y0, g_pad_y1 = g_pad
32
+
33
+ grad_output = grad_output.reshape(-1, out_size[0], out_size[1], 1)
34
+
35
+ grad_input = upfirdn2d_ext.upfirdn2d(
36
+ grad_output,
37
+ grad_kernel,
38
+ down_x,
39
+ down_y,
40
+ up_x,
41
+ up_y,
42
+ g_pad_x0,
43
+ g_pad_x1,
44
+ g_pad_y0,
45
+ g_pad_y1,
46
+ )
47
+ grad_input = grad_input.view(in_size[0], in_size[1], in_size[2], in_size[3])
48
+
49
+ ctx.save_for_backward(kernel)
50
+
51
+ pad_x0, pad_x1, pad_y0, pad_y1 = pad
52
+
53
+ ctx.up_x = up_x
54
+ ctx.up_y = up_y
55
+ ctx.down_x = down_x
56
+ ctx.down_y = down_y
57
+ ctx.pad_x0 = pad_x0
58
+ ctx.pad_x1 = pad_x1
59
+ ctx.pad_y0 = pad_y0
60
+ ctx.pad_y1 = pad_y1
61
+ ctx.in_size = in_size
62
+ ctx.out_size = out_size
63
+
64
+ return grad_input
65
+
66
+ @staticmethod
67
+ def backward(ctx, gradgrad_input):
68
+ kernel, = ctx.saved_tensors
69
+
70
+ gradgrad_input = gradgrad_input.reshape(-1, ctx.in_size[2], ctx.in_size[3], 1)
71
+
72
+ gradgrad_out = upfirdn2d_ext.upfirdn2d(
73
+ gradgrad_input,
74
+ kernel,
75
+ ctx.up_x,
76
+ ctx.up_y,
77
+ ctx.down_x,
78
+ ctx.down_y,
79
+ ctx.pad_x0,
80
+ ctx.pad_x1,
81
+ ctx.pad_y0,
82
+ ctx.pad_y1,
83
+ )
84
+ # gradgrad_out = gradgrad_out.view(ctx.in_size[0], ctx.out_size[0],
85
+ # ctx.out_size[1], ctx.in_size[3])
86
+ gradgrad_out = gradgrad_out.view(ctx.in_size[0], ctx.in_size[1], ctx.out_size[0], ctx.out_size[1])
87
+
88
+ return gradgrad_out, None, None, None, None, None, None, None, None
89
+
90
+
91
+ class UpFirDn2d(Function):
92
+
93
+ @staticmethod
94
+ def forward(ctx, input, kernel, up, down, pad):
95
+ up_x, up_y = up
96
+ down_x, down_y = down
97
+ pad_x0, pad_x1, pad_y0, pad_y1 = pad
98
+
99
+ kernel_h, kernel_w = kernel.shape
100
+ batch, channel, in_h, in_w = input.shape
101
+ ctx.in_size = input.shape
102
+
103
+ input = input.reshape(-1, in_h, in_w, 1)
104
+
105
+ ctx.save_for_backward(kernel, torch.flip(kernel, [0, 1]))
106
+
107
+ out_h = (in_h * up_y + pad_y0 + pad_y1 - kernel_h) // down_y + 1
108
+ out_w = (in_w * up_x + pad_x0 + pad_x1 - kernel_w) // down_x + 1
109
+ ctx.out_size = (out_h, out_w)
110
+
111
+ ctx.up = (up_x, up_y)
112
+ ctx.down = (down_x, down_y)
113
+ ctx.pad = (pad_x0, pad_x1, pad_y0, pad_y1)
114
+
115
+ g_pad_x0 = kernel_w - pad_x0 - 1
116
+ g_pad_y0 = kernel_h - pad_y0 - 1
117
+ g_pad_x1 = in_w * up_x - out_w * down_x + pad_x0 - up_x + 1
118
+ g_pad_y1 = in_h * up_y - out_h * down_y + pad_y0 - up_y + 1
119
+
120
+ ctx.g_pad = (g_pad_x0, g_pad_x1, g_pad_y0, g_pad_y1)
121
+
122
+ out = upfirdn2d_ext.upfirdn2d(input, kernel, up_x, up_y, down_x, down_y, pad_x0, pad_x1, pad_y0, pad_y1)
123
+ # out = out.view(major, out_h, out_w, minor)
124
+ out = out.view(-1, channel, out_h, out_w)
125
+
126
+ return out
127
+
128
+ @staticmethod
129
+ def backward(ctx, grad_output):
130
+ kernel, grad_kernel = ctx.saved_tensors
131
+
132
+ grad_input = UpFirDn2dBackward.apply(
133
+ grad_output,
134
+ kernel,
135
+ grad_kernel,
136
+ ctx.up,
137
+ ctx.down,
138
+ ctx.pad,
139
+ ctx.g_pad,
140
+ ctx.in_size,
141
+ ctx.out_size,
142
+ )
143
+
144
+ return grad_input, None, None, None, None
145
+
146
+
147
+ def upfirdn2d(input, kernel, up=1, down=1, pad=(0, 0)):
148
+ if input.device.type == 'cpu':
149
+ out = upfirdn2d_native(input, kernel, up, up, down, down, pad[0], pad[1], pad[0], pad[1])
150
+ else:
151
+ out = UpFirDn2d.apply(input, kernel, (up, up), (down, down), (pad[0], pad[1], pad[0], pad[1]))
152
+
153
+ return out
154
+
155
+
156
+ def upfirdn2d_native(input, kernel, up_x, up_y, down_x, down_y, pad_x0, pad_x1, pad_y0, pad_y1):
157
+ _, channel, in_h, in_w = input.shape
158
+ input = input.reshape(-1, in_h, in_w, 1)
159
+
160
+ _, in_h, in_w, minor = input.shape
161
+ kernel_h, kernel_w = kernel.shape
162
+
163
+ out = input.view(-1, in_h, 1, in_w, 1, minor)
164
+ out = F.pad(out, [0, 0, 0, up_x - 1, 0, 0, 0, up_y - 1])
165
+ out = out.view(-1, in_h * up_y, in_w * up_x, minor)
166
+
167
+ out = F.pad(out, [0, 0, max(pad_x0, 0), max(pad_x1, 0), max(pad_y0, 0), max(pad_y1, 0)])
168
+ out = out[:, max(-pad_y0, 0):out.shape[1] - max(-pad_y1, 0), max(-pad_x0, 0):out.shape[2] - max(-pad_x1, 0), :, ]
169
+
170
+ out = out.permute(0, 3, 1, 2)
171
+ out = out.reshape([-1, 1, in_h * up_y + pad_y0 + pad_y1, in_w * up_x + pad_x0 + pad_x1])
172
+ w = torch.flip(kernel, [0, 1]).view(1, 1, kernel_h, kernel_w)
173
+ out = F.conv2d(out, w)
174
+ out = out.reshape(
175
+ -1,
176
+ minor,
177
+ in_h * up_y + pad_y0 + pad_y1 - kernel_h + 1,
178
+ in_w * up_x + pad_x0 + pad_x1 - kernel_w + 1,
179
+ )
180
+ out = out.permute(0, 2, 3, 1)
181
+ out = out[:, ::down_y, ::down_x, :]
182
+
183
+ out_h = (in_h * up_y + pad_y0 + pad_y1 - kernel_h) // down_y + 1
184
+ out_w = (in_w * up_x + pad_x0 + pad_x1 - kernel_w) // down_x + 1
185
+
186
+ return out.view(-1, channel, out_h, out_w)