CoolFace
Apppublic

hololens/stable-diffusion-webui-depthmap-script

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
spvcnn_classsification.py161 linesDownload Raw Back to lib
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