(self, src_pc, tar_pc, key_pts, w_pc)
| 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) |
nothing calls this directly
no test coverage detected