dmfenton/splatt3r
0
1import random2 3import numpy as np4import PIL5import torch6import torchvision7 8from src.mast3r_src.dust3r.dust3r.datasets.utils.transforms import ImgNorm9from src.mast3r_src.dust3r.dust3r.utils.geometry import depthmap_to_absolute_camera_coordinates, geotrf10from src.mast3r_src.dust3r.dust3r.utils.misc import invalid_to_zeros11import src.mast3r_src.dust3r.dust3r.datasets.utils.cropping as cropping12 13 14def crop_resize_if_necessary(image, depthmap, intrinsics, resolution):15 """Adapted from DUST3R's Co3D dataset implementation"""16 17 if not isinstance(image, PIL.Image.Image):18 image = PIL.Image.fromarray(image)19 20 # Downscale with lanczos interpolation so that image.size == resolution cropping centered on the principal point21 # The new window will be a rectangle of size (2*min_margin_x, 2*min_margin_y) centered on (cx,cy)22 W, H = image.size23 cx, cy = intrinsics[:2, 2].round().astype(int)24 min_margin_x = min(cx, W - cx)25 min_margin_y = min(cy, H - cy)26 assert min_margin_x > W / 527 assert min_margin_y > H / 528 l, t = cx - min_margin_x, cy - min_margin_y29 r, b = cx + min_margin_x, cy + min_margin_y30 crop_bbox = (l, t, r, b)31 image, depthmap, intrinsics = cropping.crop_image_depthmap(image, depthmap, intrinsics, crop_bbox)32 33 # High-quality Lanczos down-scaling34 target_resolution = np.array(resolution)35 image, depthmap, intrinsics = cropping.rescale_image_depthmap(image, depthmap, intrinsics, target_resolution)36 37 # Actual cropping (if necessary) with bilinear interpolation38 intrinsics2 = cropping.camera_matrix_of_crop(intrinsics, image.size, resolution, offset_factor=0.5)39 crop_bbox = cropping.bbox_from_intrinsics_in_out(intrinsics, intrinsics2, resolution)40 image, depthmap, intrinsics2 = cropping.crop_image_depthmap(image, depthmap, intrinsics, crop_bbox)41 42 return image, depthmap, intrinsics243 44 45class DUST3RSplattingDataset(torch.utils.data.Dataset):46 47 def __init__(self, data, coverage, resolution, num_epochs_per_epoch=1, alpha=0.3, beta=0.3):48 49 super(DUST3RSplattingDataset, self).__init__()50 self.data = data51 self.coverage = coverage52 53 self.num_context_views = 254 self.num_target_views = 355 56 self.resolution = resolution57 self.transform = ImgNorm58 self.org_transform = torchvision.transforms.ToTensor()59 self.num_epochs_per_epoch = num_epochs_per_epoch60 61 self.alpha = alpha62 self.beta = beta63 64 def __getitem__(self, idx):65 66 sequence = self.data.sequences[idx // self.num_epochs_per_epoch]67 sequence_length = len(self.data.color_paths[sequence])68 69 context_views, target_views = self.sample(sequence, self.num_target_views, self.alpha, self.beta)70 71 views = {"context": [], "target": [], "scene": sequence}72 73 # Fetch the context views74 for c_view in context_views:75 76 assert c_view < sequence_length, f"Invalid view index: {c_view}, sequence length: {sequence_length}, c_views: {context_views}"77 78 view = self.data.get_view(sequence, c_view, self.resolution)79 80 # Transform the input81 view['img'] = self.transform(view['original_img'])82 view['original_img'] = self.org_transform(view['original_img'])83 84 # Create the point cloud and validity mask85 pts3d, valid_mask = depthmap_to_absolute_camera_coordinates(**view)86 view['pts3d'] = pts3d87 view['valid_mask'] = valid_mask & np.isfinite(pts3d).all(axis=-1)88 assert view['valid_mask'].any(), f"Invalid mask for sequence: {sequence}, view: {c_view}"89 90 views['context'].append(view)91 92 # Fetch the target views93 for t_view in target_views:94 95 view = self.data.get_view(sequence, t_view, self.resolution)96 view['original_img'] = self.org_transform(view['original_img'])97 views['target'].append(view)98 99 return views100 101 def __len__(self):102 103 return len(self.data.sequences) * self.num_epochs_per_epoch104 105 def sample(self, sequence, num_target_views, context_overlap_threshold=0.5, target_overlap_threshold=0.6):106 107 first_context_view = random.randint(0, len(self.data.color_paths[sequence]) - 1)108 109 # Pick a second context view that has sufficient overlap with the first context view110 valid_second_context_views = []111 for frame in range(len(self.data.color_paths[sequence])):112 if frame == first_context_view:113 continue114 overlap = self.coverage[sequence][first_context_view][frame]115 if overlap > context_overlap_threshold:116 valid_second_context_views.append(frame)117 if len(valid_second_context_views) > 0:118 second_context_view = random.choice(valid_second_context_views)119 120 # If there are no valid second context views, pick the best one121 else:122 best_view = None123 best_overlap = None124 for frame in range(len(self.data.color_paths[sequence])):125 if frame == first_context_view:126 continue127 overlap = self.coverage[sequence][first_context_view][frame]128 if best_view is None or overlap > best_overlap:129 best_view = frame130 best_overlap = overlap131 second_context_view = best_view132 133 # Pick the target views134 valid_target_views = []135 for frame in range(len(self.data.color_paths[sequence])):136 if frame == first_context_view or frame == second_context_view:137 continue138 overlap_max = max(139 self.coverage[sequence][first_context_view][frame],140 self.coverage[sequence][second_context_view][frame]141 )142 if overlap_max > target_overlap_threshold:143 valid_target_views.append(frame)144 if len(valid_target_views) >= num_target_views:145 target_views = random.sample(valid_target_views, num_target_views)146 147 # If there are not enough valid target views, pick the best ones148 else:149 overlaps = []150 for frame in range(len(self.data.color_paths[sequence])):151 if frame == first_context_view or frame == second_context_view:152 continue153 overlap = max(154 self.coverage[sequence][first_context_view][frame],155 self.coverage[sequence][second_context_view][frame]156 )157 overlaps.append((frame, overlap))158 overlaps.sort(key=lambda x: x[1], reverse=True)159 target_views = [frame for frame, _ in overlaps[:num_target_views]]160 161 return [first_context_view, second_context_view], target_views162 163 164class DUST3RSplattingTestDataset(torch.utils.data.Dataset):165 166 def __init__(self, data, samples, resolution):167 168 self.data = data169 self.samples = samples170 171 self.resolution = resolution172 self.transform = ImgNorm173 self.org_transform = torchvision.transforms.ToTensor()174 175 def get_view(self, sequence, c_view):176 177 view = self.data.get_view(sequence, c_view, self.resolution)178 179 # Transform the input180 view['img'] = self.transform(view['original_img'])181 view['original_img'] = self.org_transform(view['original_img'])182 183 # Create the point cloud and validity mask184 pts3d, valid_mask = depthmap_to_absolute_camera_coordinates(**view)185 view['pts3d'] = pts3d186 view['valid_mask'] = valid_mask & np.isfinite(pts3d).all(axis=-1)187 assert view['valid_mask'].any(), f"Invalid mask for sequence: {sequence}, view: {c_view}"188 189 return view190 191 def __getitem__(self, idx):192 193 sequence, c_view_1, c_view_2, target_view = self.samples[idx]194 c_view_1, c_view_2, target_view = int(c_view_1), int(c_view_2), int(target_view)195 fetched_c_view_1 = self.get_view(sequence, c_view_1)196 fetched_c_view_2 = self.get_view(sequence, c_view_2)197 fetched_target_view = self.get_view(sequence, target_view)198 199 views = {"context": [fetched_c_view_1, fetched_c_view_2], "target": [fetched_target_view], "scene": sequence}200 201 return views202 203 def __len__(self):204 205 return len(self.samples)206 