naver/PUMP
1
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 