hololens/stable-diffusion-webui-depthmap-script
1
1import torch.nn as nn
2
3import torchsparse.nn as spnn
4from torchsparse.point_tensor import PointTensor
5from lib.spvcnn_utils import *
6__all__ = ['SPVCNN_CLASSIFICATION']
7
8
9
10class BasicConvolutionBlock(nn.Module):
11 def __init__(self, inc, outc, ks=3, stride=1, dilation=1):
12 super().__init__()
13 self.net = nn.Sequential(
14 spnn.Conv3d(inc,
15 outc,
16 kernel_size=ks,
17 dilation=dilation,
18 stride=stride),
19 spnn.BatchNorm(outc),
20 spnn.ReLU(True))
21
22 def forward(self, x):
23 out = self.net(x)
24 return out
25
26
27class BasicDeconvolutionBlock(nn.Module):
28 def __init__(self, inc, outc, ks=3, stride=1):
29 super().__init__()
30 self.net = nn.Sequential(
31 spnn.Conv3d(inc,
32 outc,
33 kernel_size=ks,
34 stride=stride,
35 transpose=True),
36 spnn.BatchNorm(outc),
37 spnn.ReLU(True))
38
39 def forward(self, x):
40 return self.net(x)
41
42
43class ResidualBlock(nn.Module):
44 def __init__(self, inc, outc, ks=3, stride=1, dilation=1):
45 super().__init__()
46 self.net = nn.Sequential(
47 spnn.Conv3d(inc,
48 outc,
49 kernel_size=ks,
50 dilation=dilation,
51 stride=stride), spnn.BatchNorm(outc),
52 spnn.ReLU(True),
53 spnn.Conv3d(outc,
54 outc,
55 kernel_size=ks,
56 dilation=dilation,
57 stride=1),
58 spnn.BatchNorm(outc)
59 )
60
61 self.downsample = nn.Sequential() if (inc == outc and stride == 1) else \
62 nn.Sequential(
63 spnn.Conv3d(inc, outc, kernel_size=1, dilation=1, stride=stride),
64 spnn.BatchNorm(outc)
65 )
66
67 self.relu = spnn.ReLU(True)
68
69 def forward(self, x):
70 out = self.relu(self.net(x) + self.downsample(x))
71 return out
72
73
74class SPVCNN_CLASSIFICATION(nn.Module):
75 def __init__(self, **kwargs):
76 super().__init__()
77
78 cr = kwargs.get('cr', 1.0)
79 cs = [32, 32, 64, 128, 256, 256, 128, 96, 96]
80 cs = [int(cr * x) for x in cs]
81
82 if 'pres' in kwargs and 'vres' in kwargs:
83 self.pres = kwargs['pres']
84 self.vres = kwargs['vres']
85
86 self.stem = nn.Sequential(
87 spnn.Conv3d(kwargs['input_channel'], cs[0], kernel_size=3, stride=1),
88 spnn.BatchNorm(cs[0]),
89 spnn.ReLU(True),
90 spnn.Conv3d(cs[0], cs[0], kernel_size=3, stride=1),
91 spnn.BatchNorm(cs[0]),
92 spnn.ReLU(True))
93
94 self.stage1 = nn.Sequential(
95 BasicConvolutionBlock(cs[0], cs[0], ks=2, stride=2, dilation=1),
96 ResidualBlock(cs[0], cs[1], ks=3, stride=1, dilation=1),
97 ResidualBlock(cs[1], cs[1], ks=3, stride=1, dilation=1),
98 )
99
100 self.stage2 = nn.Sequential(
101 BasicConvolutionBlock(cs[1], cs[1], ks=2, stride=2, dilation=1),
102 ResidualBlock(cs[1], cs[2], ks=3, stride=1, dilation=1),
103 ResidualBlock(cs[2], cs[2], ks=3, stride=1, dilation=1),
104 )
105
106 self.stage3 = nn.Sequential(
107 BasicConvolutionBlock(cs[2], cs[2], ks=2, stride=2, dilation=1),
108 ResidualBlock(cs[2], cs[3], ks=3, stride=1, dilation=1),
109 ResidualBlock(cs[3], cs[3], ks=3, stride=1, dilation=1),
110 )
111
112 self.stage4 = nn.Sequential(
113 BasicConvolutionBlock(cs[3], cs[3], ks=2, stride=2, dilation=1),
114 ResidualBlock(cs[3], cs[4], ks=3, stride=1, dilation=1),
115 ResidualBlock(cs[4], cs[4], ks=3, stride=1, dilation=1),
116 )
117 self.avg_pool = spnn.GlobalAveragePooling()
118 self.classifier = nn.Sequential(nn.Linear(cs[4], kwargs['num_classes']))
119 self.point_transforms = nn.ModuleList([
120 nn.Sequential(
121 nn.Linear(cs[0], cs[4]),
122 nn.BatchNorm1d(cs[4]),
123 nn.ReLU(True),
124 ),
125 ])
126
127 self.weight_initialization()
128 self.dropout = nn.Dropout(0.3, True)
129
130 def weight_initialization(self):
131 for m in self.modules():
132 if isinstance(m, nn.BatchNorm1d):
133 nn.init.constant_(m.weight, 1)
134 nn.init.constant_(m.bias, 0)
135
136 def forward(self, x):
137 # x: SparseTensor z: PointTensor
138 z = PointTensor(x.F, x.C.float())
139
140 x0 = initial_voxelize(z, self.pres, self.vres)
141
142 x0 = self.stem(x0)
143 z0 = voxel_to_point(x0, z, nearest=False)
144 z0.F = z0.F
145
146 x1 = point_to_voxel(x0, z0)
147 x1 = self.stage1(x1)
148 x2 = self.stage2(x1)
149 x3 = self.stage3(x2)
150 x4 = self.stage4(x3)
151 z1 = voxel_to_point(x4, z0)
152 z1.F = z1.F + self.point_transforms[0](z0.F)
153 y1 = point_to_voxel(x4, z1)
154 pool = self.avg_pool(y1)
155 out = self.classifier(pool)
156
157
158 return out
159
160
161 