(self)
| 289 | |
| 290 | @torch.no_grad() |
| 291 | def evaluate(self): |
| 292 | self.net.eval() |
| 293 | |
| 294 | feat_maps = self.encode() |
| 295 | sdf_grid_pred = self.decode_batch(feat_maps, self.pts_grid)[..., :1] |
| 296 | |
| 297 | if self.sdf_renorm: |
| 298 | sdf_grid_pred = sdf_grid_pred * self.sdf_threshold |
| 299 | sdf_grid_gt = self.sdf_grid * self.sdf_threshold |
| 300 | else: |
| 301 | sdf_grid_gt = self.sdf_grid |
| 302 | stat = evaluate_tsdf_prediction(sdf_grid_pred, sdf_grid_gt, self.sdf_threshold) |
| 303 | |
| 304 | if self.data_type != "sdf": |
| 305 | tex_surf_pred = self.decode_batch(feat_maps, self.pts_on_surf)[..., 1:] # (n, 3) |
| 306 | surf_tex_l1_error = (tex_surf_pred - self.tex_on_surf).abs().mean().item() |
| 307 | stat["surf_tex_l1_error"] = surf_tex_l1_error |
| 308 | |
| 309 | return stat |
| 310 | |
| 311 | @torch.no_grad() |
| 312 | def encode(self, vol=None): |
no test coverage detected