radames/Text2Human-API
1
1import os2import os.path3 4import numpy as np5import torch6import torch.utils.data as data7from PIL import Image8 9 10class ParsingGenerationDeepFashionAttrSegmDataset(data.Dataset):11 12 def __init__(self, segm_dir, pose_dir, ann_file, downsample_factor=2):13 self._densepose_path = pose_dir14 self._segm_path = segm_dir15 self._image_fnames = []16 self.attrs = []17 18 self.downsample_factor = downsample_factor19 20 # training, ground-truth available21 assert os.path.exists(ann_file)22 for row in open(os.path.join(ann_file), 'r'):23 annotations = row.split()24 self._image_fnames.append(annotations[0])25 self.attrs.append([int(i) for i in annotations[1:]])26 27 def _open_file(self, path_prefix, fname):28 return open(os.path.join(path_prefix, fname), 'rb')29 30 def _load_densepose(self, raw_idx):31 fname = self._image_fnames[raw_idx]32 fname = f'{fname[:-4]}_densepose.png'33 with self._open_file(self._densepose_path, fname) as f:34 densepose = Image.open(f)35 if self.downsample_factor != 1:36 width, height = densepose.size37 width = width // self.downsample_factor38 height = height // self.downsample_factor39 densepose = densepose.resize(40 size=(width, height), resample=Image.NEAREST)41 # channel-wise IUV order, [3, H, W]42 densepose = np.array(densepose)[:, :, 2:].transpose(2, 0, 1)43 return densepose.astype(np.float32)44 45 def _load_segm(self, raw_idx):46 fname = self._image_fnames[raw_idx]47 fname = f'{fname[:-4]}_segm.png'48 with self._open_file(self._segm_path, fname) as f:49 segm = Image.open(f)50 if self.downsample_factor != 1:51 width, height = segm.size52 width = width // self.downsample_factor53 height = height // self.downsample_factor54 segm = segm.resize(55 size=(width, height), resample=Image.NEAREST)56 segm = np.array(segm)57 return segm.astype(np.float32)58 59 def __getitem__(self, index):60 pose = self._load_densepose(index)61 segm = self._load_segm(index)62 attr = self.attrs[index]63 64 pose = torch.from_numpy(pose)65 segm = torch.LongTensor(segm)66 attr = torch.LongTensor(attr)67 68 pose = pose / 12. - 169 70 return_dict = {71 'densepose': pose,72 'segm': segm,73 'attr': attr,74 'img_name': self._image_fnames[index]75 }76 77 return return_dict78 79 def __len__(self):80 return len(self._image_fnames)81 