CoolFace
Apppublic

dmfenton/splatt3r

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
data.py206 linesDownload Raw Back to data
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