MCPcopy Create free account
hub / github.com/Colin97/DeepMetaHandles / forward

Method forward

src/network.py:92–180  ·  view source on GitHub ↗
(self, src_pc, tar_pc, key_pts, w_pc)

Source from the content-addressed store, hash-verified

90 self.sigmoid = nn.Sigmoid()
91
92 def forward(self, src_pc, tar_pc, key_pts, w_pc):
93 B, N, _ = src_pc.shape
94 src_out, src_global = self.pointnet(src_pc, False)
95 tar_global = self.pointnet(tar_pc, True)
96
97 src_out = F.relu(self.bn11(self.conv11(src_out)))
98 src_out = F.relu(self.bn12(self.conv12(src_out)))
99 src_out = F.relu(self.bn13(self.conv13(src_out)))
100
101 _, K, _ = key_pts.shape
102 key_pts1 = key_pts.unsqueeze(-1).expand(-1, -1, -1, N) # B K 3 N
103 w_pc1 = w_pc.transpose(2, 1).unsqueeze(2) # B K 1 N
104 src_out = src_out.unsqueeze(1).expand(-1, K, -1, -1) # B K 64 N
105 net = torch.cat([src_out, w_pc1, key_pts1], 2).view(B * K, 68, N)
106
107 net = F.relu(self.bn21(self.conv21(net)))
108 net = self.bn22(self.conv22(net))
109
110 net = torch.max(net, 2, keepdim=True)[0]
111 key_fea = net.view(B * K, 64, 1)
112
113 net = torch.cat([key_fea, key_pts.view(B * K, 3, 1)], 1)
114 net = F.relu(self.bn31(self.conv31(net)))
115 net = F.relu(self.bn32(self.conv32(net)))
116 basis = self.conv33(net).view(B, K * 3, self.num_basis).transpose(1, 2)
117 basis = basis / basis.norm(p=2, dim=-1, keepdim=True)
118
119 key_fea_range = key_fea.view(
120 B, K, 64, 1).expand(-1, -1, -1, self.num_basis).transpose(1, 3)
121 key_pts_range = key_pts.view(
122 B, K, 3, 1).expand(-1, -1, -1, self.num_basis).transpose(1, 3)
123 basis_range = basis.view(B, self.num_basis, K, 3).transpose(2, 3)
124
125 coef_range = torch.cat([key_fea_range, key_pts_range, basis_range], 2).view(
126 B * self.num_basis, 70, K)
127 coef_range = F.relu(self.bn71(self.conv71(coef_range)))
128 coef_range = F.relu(self.bn72(self.conv72(coef_range)))
129 coef_range = self.conv73(coef_range)
130 coef_range = torch.max(coef_range, 2, keepdim=True)[0]
131 coef_range = coef_range.view(B, self.num_basis, 2) * 0.01
132 coef_range[:, :, 0] = coef_range[:, :, 0] * -1
133
134 src_tar = torch.cat([src_global, tar_global], 1).unsqueeze(
135 1).expand(-1, K, -1).reshape(B * K, 2048, 1)
136
137 key_fea = torch.cat([key_fea, src_tar, key_pts.view(B * K, 3, 1)], 1)
138 key_fea = F.relu(self.bn41(self.conv41(key_fea)))
139 key_fea = F.relu(self.bn42(self.conv42(key_fea)))
140 key_fea = F.relu(self.bn43(self.conv43(key_fea)))
141
142 key_fea = key_fea.view(B, K, 128).transpose(
143 1, 2).unsqueeze(1).expand(-1, self.num_basis, -1, -1)
144 key_pts2 = key_pts.view(B, K, 3).transpose(
145 1, 2).unsqueeze(1).expand(-1, self.num_basis, -1, -1)
146 basis1 = basis.view(B, self.num_basis, K, 3).transpose(2, 3)
147
148 net = torch.cat([key_fea, basis1, key_pts2], 2).view(
149 B * self.num_basis, 3 + 128 + 3, K)

Callers

nothing calls this directly

Calls 1

chamfer_distanceFunction · 0.90

Tested by

no test coverage detected