(self, x, feat_maps, aabb=None)
| 87 | return feat |
| 88 | |
| 89 | def decode(self, x, feat_maps, aabb=None): |
| 90 | # x [N, 3] |
| 91 | if aabb is None: |
| 92 | aabb = self.aabb |
| 93 | x = 2 * (x - aabb[:3]) / (aabb[3:] - aabb[:3]) - 1 # [-1, 1] |
| 94 | |
| 95 | h_geo = 0 |
| 96 | h_tex = 0 |
| 97 | |
| 98 | coords_list = [[0, 1], [0, 2], [1, 2]] |
| 99 | |
| 100 | geo_feat_maps = [fm[:, :self.geo_feat_dim] for fm in feat_maps] |
| 101 | geo_feat_maps = self.geo_convs(geo_feat_maps) |
| 102 | for i in range(3): |
| 103 | h_geo += self.sample_feature_plane2D(geo_feat_maps[i], x[..., coords_list[i]]) # (N, C) |
| 104 | |
| 105 | if self.use_tex: |
| 106 | tex_feat_maps = [fm[:, self.geo_feat_dim:] for fm in feat_maps] |
| 107 | tex_feat_maps = self.tex_convs(tex_feat_maps) |
| 108 | for i in range(3): |
| 109 | h_tex += self.sample_feature_plane2D(tex_feat_maps[i], x[..., coords_list[i]]) |
| 110 | |
| 111 | h_geo = self.geo_decoder(h_geo) # (N, 1) |
| 112 | if self.use_tex: |
| 113 | h_tex = self.tex_decoder(h_tex).sigmoid() # (N, 1) |
| 114 | h = torch.cat([h_geo, h_tex], dim=1) |
| 115 | else: |
| 116 | h = h_geo |
| 117 | return h |
| 118 | |
| 119 | def forward(self, vol, x, aabb=None): |
| 120 | feat_map = self.encode(vol) |
no test coverage detected