Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
parsing_generation_segm_attr_dataset.py81 linesDownload Raw Back to data
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