Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
pose_attr_dataset.py110 linesDownload Raw Back to data
1import os2import os.path3import random4 5import numpy as np6import torch7import torch.utils.data as data8from PIL import Image9 10 11class DeepFashionAttrPoseDataset(data.Dataset):12 13    def __init__(self,14                 pose_dir,15                 texture_ann_dir,16                 shape_ann_path,17                 downsample_factor=2,18                 xflip=False):19        self._densepose_path = pose_dir20        self._image_fnames_target = []21        self._image_fnames = []22        self.upper_fused_attrs = []23        self.lower_fused_attrs = []24        self.outer_fused_attrs = []25        self.shape_attrs = []26 27        self.downsample_factor = downsample_factor28        self.xflip = xflip29 30        # load attributes31        assert os.path.exists(f'{texture_ann_dir}/upper_fused.txt')32        for idx, row in enumerate(33                open(os.path.join(f'{texture_ann_dir}/upper_fused.txt'), 'r')):34            annotations = row.split()35            self._image_fnames_target.append(annotations[0])36            self._image_fnames.append(f'{annotations[0].split(".")[0]}.png')37            self.upper_fused_attrs.append(int(annotations[1]))38 39        assert len(self._image_fnames_target) == len(self.upper_fused_attrs)40 41        assert os.path.exists(f'{texture_ann_dir}/lower_fused.txt')42        for idx, row in enumerate(43                open(os.path.join(f'{texture_ann_dir}/lower_fused.txt'), 'r')):44            annotations = row.split()45            assert self._image_fnames_target[idx] == annotations[0]46            self.lower_fused_attrs.append(int(annotations[1]))47 48        assert len(self._image_fnames_target) == len(self.lower_fused_attrs)49 50        assert os.path.exists(f'{texture_ann_dir}/outer_fused.txt')51        for idx, row in enumerate(52                open(os.path.join(f'{texture_ann_dir}/outer_fused.txt'), 'r')):53            annotations = row.split()54            assert self._image_fnames_target[idx] == annotations[0]55            self.outer_fused_attrs.append(int(annotations[1]))56 57        assert len(self._image_fnames_target) == len(self.outer_fused_attrs)58 59        assert os.path.exists(shape_ann_path)60        for idx, row in enumerate(open(os.path.join(shape_ann_path), 'r')):61            annotations = row.split()62            assert self._image_fnames_target[idx] == annotations[0]63            self.shape_attrs.append([int(i) for i in annotations[1:]])64 65    def _open_file(self, path_prefix, fname):66        return open(os.path.join(path_prefix, fname), 'rb')67 68    def _load_densepose(self, raw_idx):69        fname = self._image_fnames[raw_idx]70        fname = f'{fname[:-4]}_densepose.png'71        with self._open_file(self._densepose_path, fname) as f:72            densepose = Image.open(f)73            if self.downsample_factor != 1:74                width, height = densepose.size75                width = width // self.downsample_factor76                height = height // self.downsample_factor77                densepose = densepose.resize(78                    size=(width, height), resample=Image.NEAREST)79            # channel-wise IUV order, [3, H, W]80            densepose = np.array(densepose)[:, :, 2:].transpose(2, 0, 1)81        return densepose.astype(np.float32)82 83    def __getitem__(self, index):84        pose = self._load_densepose(index)85        shape_attr = self.shape_attrs[index]86        shape_attr = torch.LongTensor(shape_attr)87 88        if self.xflip and random.random() > 0.5:89            pose = pose[:, :, ::-1].copy()90 91        upper_fused_attr = self.upper_fused_attrs[index]92        lower_fused_attr = self.lower_fused_attrs[index]93        outer_fused_attr = self.outer_fused_attrs[index]94 95        pose = pose / 12. - 196 97        return_dict = {98            'densepose': pose,99            'img_name': self._image_fnames_target[index],100            'shape_attr': shape_attr,101            'upper_fused_attr': upper_fused_attr,102            'lower_fused_attr': lower_fused_attr,103            'outer_fused_attr': outer_fused_attr,104        }105 106        return return_dict107 108    def __len__(self):109        return len(self._image_fnames)110