-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathHLKConv.py
More file actions
40 lines (33 loc) · 1.91 KB
/
Copy pathHLKConv.py
File metadata and controls
40 lines (33 loc) · 1.91 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
import torch
import torch.nn as nn
class HLKConv(nn.Module):
def __init__(self, dim, k_size):
super().__init__()
self.k_size = k_size
if k_size == 7:
self.conv0 = nn.Conv2d(dim, dim, kernel_size=3, stride=1, padding=((3-1)//2), groups=dim)
self.conv_spatial = nn.Conv2d(dim, dim, kernel_size=3, stride=1, padding=2, groups=dim, dilation=2)
elif k_size == 11:
self.conv0 = nn.Conv2d(dim, dim, kernel_size=3, stride=1, padding=((3-1)//2), groups=dim)
self.conv_spatial = nn.Conv2d(dim, dim, kernel_size=5, stride=1, padding=4, groups=dim, dilation=2)
elif k_size == 23:
self.conv0 = nn.Conv2d(dim, dim, kernel_size=5, stride=1, padding=((5-1)//2), groups=dim)
self.conv_spatial = nn.Conv2d(dim, dim, kernel_size=7, stride=1, padding=9, groups=dim, dilation=3)
elif k_size == 35:
self.conv0 = nn.Conv2d(dim, dim, kernel_size=5, stride=1, padding=((5-1)//2), groups=dim)
self.conv_spatial = nn.Conv2d(dim, dim, kernel_size=11, stride=1, padding=15, groups=dim, dilation=3)
elif k_size == 41:
self.conv0 = nn.Conv2d(dim, dim, kernel_size=5, stride=1, padding=((5-1)//2), groups=dim)
self.conv_spatial = nn.Conv2d(dim, dim, kernel_size=13, stride=1, padding=18, groups=dim, dilation=3)
elif k_size == 53:
self.conv0 = nn.Conv2d(dim, dim, kernel_size=5, stride=1, padding=((5-1)//2), groups=dim)
self.conv_spatial = nn.Conv2d(dim, dim, kernel_size=17, stride=1, padding=24, groups=dim, dilation=3)
self.conv1 = nn.Conv2d(dim * 2, dim, 1)
def forward(self, x):
x = self.conv0(x)
x = self.conv1(torch.cat([x, self.conv_spatial(x)], dim=1))
return x
if __name__ == '__main__':
x = torch.randn((2, 2, 512, 512))
model = HLKConv(2, 53)
print(model(x).shape)