(self, decoder, map_states, voxel_size, voxels,
frame_poses=None, depth_maps=None, clean_mseh=False,
require_color=False, offset=-10, res=8)
| 77 | |
| 78 | @torch.no_grad() |
| 79 | def create_mesh(self, decoder, map_states, voxel_size, voxels, |
| 80 | frame_poses=None, depth_maps=None, clean_mseh=False, |
| 81 | require_color=False, offset=-10, res=8): |
| 82 | |
| 83 | sdf_grid = get_scores(decoder, map_states, voxel_size, bits=res, model=self.model) # (num_voxels,res*3,4) |
| 84 | |
| 85 | sdf_grid = sdf_grid.reshape(-1, res, res, res, 1) |
| 86 | |
| 87 | voxel_centres = map_states["voxel_center_xyz"] |
| 88 | verts, faces = self.marching_cubes(voxel_centres, sdf_grid) |
| 89 | |
| 90 | if clean_mseh: |
| 91 | print("********** get points from frames **********") |
| 92 | all_points = self.get_valid_points(frame_poses, depth_maps) |
| 93 | print("********** construct kdtree **********") |
| 94 | kdtree = cKDTree(all_points) |
| 95 | print("********** query kdtree **********") |
| 96 | point_mask = kdtree.query_ball_point( |
| 97 | verts, voxel_size * 0.5, workers=12, return_length=True) |
| 98 | print("********** finished querying kdtree **********") |
| 99 | point_mask = point_mask > 0 |
| 100 | face_mask = point_mask[faces.reshape(-1)].reshape(-1, 3).any(-1) |
| 101 | |
| 102 | faces = faces[face_mask] |
| 103 | |
| 104 | if require_color: |
| 105 | print("********** get color from network **********") |
| 106 | verts_torch = torch.from_numpy(verts).float().cuda() |
| 107 | batch_points = torch.split(verts_torch, 1000) |
| 108 | colors = [] |
| 109 | for points in batch_points: |
| 110 | # voxel_pos = points // self.voxel_size |
| 111 | voxel_pos = torch.div(points, self.voxel_size, rounding_mode='trunc') |
| 112 | batch_voxels = voxels[:, :3].cuda() |
| 113 | batch_voxels = batch_voxels.unsqueeze( |
| 114 | 0).repeat(voxel_pos.shape[0], 1, 1) |
| 115 | |
| 116 | # filter outliers |
| 117 | nonzeros = (batch_voxels == voxel_pos.unsqueeze(1)).all(-1) |
| 118 | nonzeros = torch.where(nonzeros, torch.ones_like( |
| 119 | nonzeros).int(), -torch.ones_like(nonzeros).int()) |
| 120 | sorted, index = torch.sort(nonzeros, dim=-1, descending=True) |
| 121 | sorted = sorted[:, 0] |
| 122 | index = index[:, 0] |
| 123 | valid = (sorted != -1) |
| 124 | color_empty = torch.zeros_like(points) |
| 125 | points = points[valid, :] |
| 126 | index = index[valid] |
| 127 | |
| 128 | # get color |
| 129 | if len(points) > 0: |
| 130 | color = eval_points(decoder, points).cuda() |
| 131 | color_empty[valid] = color.float() |
| 132 | colors += [color_empty] |
| 133 | colors = torch.cat(colors, 0) |
| 134 | |
| 135 | mesh = o3d.geometry.TriangleMesh() |
| 136 | mesh.vertices = o3d.utility.Vector3dVector(verts + offset) |
no test coverage detected