x: point features of shape [B, N_feat, 3, N_samples, ...]
(self, x)
| 201 | self.vn_lin = nn.Linear(in_channels // 4, 3, bias=False) |
| 202 | |
| 203 | def forward(self, x): |
| 204 | ''' |
| 205 | x: point features of shape [B, N_feat, 3, N_samples, ...] |
| 206 | ''' |
| 207 | z0 = x |
| 208 | z0 = self.vn1(z0) |
| 209 | z0 = self.vn2(z0) |
| 210 | z0 = self.vn_lin(z0.transpose(1, -1)).transpose(1, -1) |
| 211 | |
| 212 | if self.normalize_frame: |
| 213 | # make z0 orthogonal. u2 = v2 - proj_u1(v2) |
| 214 | v1 = z0[:, 0, :] |
| 215 | # u1 = F.normalize(v1, dim=1) |
| 216 | v1_norm = torch.sqrt((v1 * v1).sum(1, keepdims=True)) |
| 217 | u1 = v1 / (v1_norm + EPS) |
| 218 | v2 = z0[:, 1, :] |
| 219 | v2 = v2 - (v2 * u1).sum(1, keepdims=True) * u1 |
| 220 | # u2 = F.normalize(u2, dim=1) |
| 221 | v2_norm = torch.sqrt((v2 * v2).sum(1, keepdims=True)) |
| 222 | u2 = v2 / (v2_norm + EPS) |
| 223 | |
| 224 | # compute the cross product of the two output vectors |
| 225 | u3 = torch.cross(u1, u2) |
| 226 | z0 = torch.stack([u1, u2, u3], dim=1).transpose(1, 2) |
| 227 | else: |
| 228 | z0 = z0.transpose(1, 2) |
| 229 | |
| 230 | if self.dim == 4: |
| 231 | x_std = torch.einsum('bijm,bjkm->bikm', x, z0) |
| 232 | elif self.dim == 3: |
| 233 | x_std = torch.einsum('bij,bjk->bik', x, z0) |
| 234 | elif self.dim == 5: |
| 235 | x_std = torch.einsum('bijmn,bjkmn->bikmn', x, z0) |
| 236 | |
| 237 | return x_std, z0 |
| 238 | |
| 239 | |
| 240 | class VNInFeature(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected