CoolFace
Apppublic

naver/PUMP

sourceHugging Faceupdated 4y agoView on Hugging Face
1likes
pair_dataset.py227 linesDownload Raw Back to datasets
1# Copyright 2022-present NAVER Corp.2# CC BY-NC-SA 4.03# Available only for non-commercial use4 5from pdb import set_trace as bb6import os, os.path as osp7from tqdm import tqdm8from PIL import Image9import numpy as np10import torch11 12from .image_set import ImageSet13from .transforms import instanciate_transforms14from .utils import DatasetWithRng15invh = np.linalg.inv16 17 18class ImagePairs (DatasetWithRng):19    """ Base class for a dataset that serves image pairs.20    """21    imgs = None # regular image dataset22    pairs = [] # list of (idx1, idx2), ...23 24    def __init__(self, image_set, pairs, trf=None, **rng):25        assert image_set and pairs, 'empty images or pairs'26        super().__init__(**rng)27        self.imgs = image_set28        self.pairs = pairs29        self.trf = instanciate_transforms(trf, rng=self.rng)30 31    def __len__(self):32        return len(self.pairs)33 34    def __getitem__(self, idx):35        transform = self.trf or (lambda x:x)36        pair = tuple(map(transform, self._load_pair(idx)))37        return pair, {}38 39    def _load_pair(self, idx):40        i,j = self.pairs[idx]41        img1 = self.imgs.get_image(i)42        return (img1, img1) if i == j else (img1, self.imgs.get_image(j))43 44    def __repr__(self):45        return f'{self.__class__.__name__}({len(self)} pairs from {self.imgs})'46 47 48class StillImagePairs (ImagePairs):49    """ A dataset of 'still' image pairs used for debugging purposes.50    """51    def __init__(self, image_set, pairs=None, **rng):52        if isinstance(image_set, ImagePairs):53            super().__init__(image_set.imgs, pairs or image_set.pairs, **rng)54        else:55            super().__init__(image_set, pairs or [(i,i) for i in range(len(image_set))], **rng)56 57    def __getitem__(self, idx):58        img1, img2 = self._load_pair(idx)59        sx, sy = img2.size / np.float32(img1.size)60        return (img1, img2), dict(homography=np.diag(np.float32([sx, sy, 1])))61 62 63class SyntheticImagePairs (StillImagePairs):64    """ A synthetic generator of image pairs.65        Given a normal image dataset, it constructs pairs using random homographies & noise.66 67    scale: prior image scaling.68    distort: distortion applied independently to (img1,img2) if sym=True else just img269    sym: (bool) see above.70    """71    def __init__(self, image_set, scale='', distort='', sym=False, **rng):72        super().__init__(image_set, **rng)73        self.symmetric = sym74        self.scale = instanciate_transforms(scale, rng=self.rng)75        self.distort = instanciate_transforms(distort, rng=self.rng)76 77    def __getitem__(self, idx):78        (img1, img2), gt = super().__getitem__(idx)79 80        img1 = dict(img=img1, homography=np.eye(3,dtype=np.float32))81        if img1['img'] is img2:82            img1 = self.scale(img1)83            img2 = self.distort(dict(img1))84            if self.symmetric: img1 = self.distort(img1)85        else:86            if self.symmetric: img1 = self.distort(self.scale(img1))87            img2 = self.distort(self.scale(dict(img=img2, **gt)))88 89        return (img1['img'], img2['img']), dict(homography=img2['homography'] @ invh(img1['homography']))90 91    def __repr__(self):92        format = lambda s: ','.join(l.strip() for l in repr(s).splitlines() if l).replace(',','',1)93        return f"{self.__class__.__name__}({len(self)} images, scale={format(self.scale)}, distort={format(self.distort)})"94 95 96class CatImagePairs (DatasetWithRng):97    """ Concatenation of several ImagePairs datasets98    """99    def __init__(self, *pair_datasets, seed=torch.initial_seed()):100        assert all(isinstance(db, ImagePairs) for db in pair_datasets)101        self.pair_datasets = pair_datasets102        DatasetWithRng.__init__(self, seed=seed) # init last103        self._init()104 105    def _init(self):106        self._pair_offsets = np.cumsum([0] + [len(db) for db in self.pair_datasets])107        self.npairs = self._pair_offsets[-1]108 109    def __len__(self):110        return self.npairs111 112    def __repr__(self):113        fmt_str = f"{type(self).__name__}({len(self)} pairs,"114        for i,db in enumerate(self.pair_datasets):115            npairs = self._pair_offsets[i+1] - self._pair_offsets[i]116            fmt_str += f'\n\t{npairs} from '+str(db).replace("\n"," ") + ','117        return fmt_str[:-1] + ')'118 119    def __getitem__(self, idx):120        b, i = self._which(idx)121        return self.pair_datasets[b].__getitem__(i)122 123    def _which(self, i):124        pos = np.searchsorted(self._pair_offsets, i, side='right')-1125        assert pos < self.npairs, 'Bad pair index %d >= %d' % (i, self.npairs)126        return pos, i - self._pair_offsets[pos]127 128    def _call(self, func, i, *args, **kwargs):129        b, j = self._which(i)130        return getattr(self.pair_datasets[b], func)(j, *args, **kwargs)131 132    def init_worker(self, tid):133        for db in self.pair_datasets:134            db.init_worker(tid)135 136 137class BalancedCatImagePairs (CatImagePairs):138    """ Balanced concatenation of several ImagePairs datasets139    """140    def __init__(self, npairs=0, *pair_datasets, **kw):141        assert isinstance(npairs, int) and npairs >= 0, 'BalancedCatImagePairs(npairs != int)'142        assert len(pair_datasets) > 0, 'no dataset provided'143 144        if len(pair_datasets) >= 3 and isinstance(pair_datasets[1], int):145            assert len(pair_datasets) % 2 == 1146            pair_datasets = [npairs] + list(pair_datasets)147            npairs, pair_datasets = pair_datasets[0::2], pair_datasets[1::2]148            assert all(isinstance(n, int) for n in npairs)149            self._pair_offsets = np.cumsum([0]+npairs)150            self.npairs = self._pair_offsets[-1]151        else:152            self.npairs = npairs or max(len(db) for db in pair_datasets)153            self._pair_offsets = np.linspace(0, self.npairs, len(pair_datasets)+1).astype(int)154        CatImagePairs.__init__(self, *pair_datasets, **kw)155 156    def set_epoch(self, epoch):157        DatasetWithRng.init_worker(self, epoch) # random seed only depends on the epoch158        self._init() # reset permutations for this epoch159 160    def init_worker(self, tid):161        CatImagePairs.init_worker(self, tid) 162 163    def _init(self):164        self._perms = []165        for i,db in enumerate(self.pair_datasets):166            assert len(db), 'cannot balance if there is an empty dataset'167            avail = self._pair_offsets[i+1] - self._pair_offsets[i]168            idxs = np.arange(len(db))169            while len(idxs) < avail: 170                idxs = np.r_[idxs,idxs]171            if self.seed: # if not seed, then no shuffle172                self.rng.shuffle(idxs[(avail//len(db))*len(db):])173            self._perms.append( idxs[:avail] )174        # print(self._perms)175 176    def _which(self, i):177        pos, idx = super()._which(i)178        return pos, self._perms[pos][idx]179 180 181class UnsupervisedPairs (ImagePairs):182    """ Unsupervised image pairs obtained from SfM183    """184    def __init__(self, img_set, pair_file_path):185        assert isinstance(img_set, ImageSet), bb()186        self.pair_list = self._parse_pair_list(pair_file_path)187        self.corres_dir = osp.join(osp.split(pair_file_path)[0], 'corres')188 189        tag_to_idx = {n:i for i,n in enumerate(img_set.imgs)}190        img_indices = lambda pair: tuple([tag_to_idx[n] for n in pair])191        super().__init__(img_set, [img_indices(pair) for pair in self.pair_list])192 193    def __repr__(self):194        return f"{type(self).__name__}({len(self)} pairs from {self.imgs})"195 196    def _parse_pair_list(self, pair_file_path):197        res = []198        for row in open(pair_file_path).read().splitlines():199            row = row.split()200            if len(row) != 2: raise IOError()201            res.append((row[0], row[1]))202        return res203 204    def get_corres_path(self, pair_idx):205        img1, img2 = [osp.basename(self.imgs.imgs[i]) for i in self.pairs[pair_idx]]206        return osp.join(self.corres_dir, f'{img1}_{img2}.npy')207 208    def get_corres(self, pair_idx):209        return np.load(self.get_corres_path(pair_idx))210 211    def __getitem__(self, idx):212        img1, img2 = self._load_pair(idx)213        return (img1, img2), dict(corres=self.get_corres(idx))214 215 216if __name__ == '__main__':217    from datasets import *218    from tools.viz import show_random_pairs219 220    db = BalancedCatImagePairs(221                3125, SyntheticImagePairs(RandomWebImages(0,52),distort='RandomTilting(0.5)'),222                4875, SyntheticImagePairs(SfM120k_Images(),distort='RandomTilting(0.5)'),223                8000, SfM120k_Pairs())224 225    show_random_pairs(db)226    227