Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
mask_dataset.py60 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 MaskDataset(data.Dataset):12 13    def __init__(self, segm_dir, ann_dir, downsample_factor=2, xflip=False):14 15        self._segm_path = segm_dir16        self._image_fnames = []17 18        self.downsample_factor = downsample_factor19        self.xflip = xflip20 21        # load attributes22        assert os.path.exists(f'{ann_dir}/upper_fused.txt')23        for idx, row in enumerate(24                open(os.path.join(f'{ann_dir}/upper_fused.txt'), 'r')):25            annotations = row.split()26            self._image_fnames.append(annotations[0])27 28    def _open_file(self, path_prefix, fname):29        return open(os.path.join(path_prefix, fname), 'rb')30 31    def _load_segm(self, raw_idx):32        fname = self._image_fnames[raw_idx]33        fname = f'{fname[:-4]}_segm.png'34        with self._open_file(self._segm_path, fname) as f:35            segm = Image.open(f)36            if self.downsample_factor != 1:37                width, height = segm.size38                width = width // self.downsample_factor39                height = height // self.downsample_factor40                segm = segm.resize(41                    size=(width, height), resample=Image.NEAREST)42            segm = np.array(segm)43        # segm = segm[:, :, np.newaxis].transpose(2, 0, 1)44        return segm.astype(np.float32)45 46    def __getitem__(self, index):47        segm = self._load_segm(index)48 49        if self.xflip and random.random() > 0.5:50            segm = segm[:, ::-1].copy()51 52        segm = torch.from_numpy(segm).long()53 54        return_dict = {'segm': segm, 'img_name': self._image_fnames[index]}55 56        return return_dict57 58    def __len__(self):59        return len(self._image_fnames)60