CoolFace
Apppublic

KyanChen/BuildingExtraction

sourceHugging Faceupdated 4y agoView on Hugging Face
5likes
Datasets.py144 linesDownload Raw Back to Utils
1import os.path2 3from torch.utils.data import Dataset, DataLoader4import torch5import numpy as np6import pandas as pd7from skimage import io8from Utils.Augmentations import Augmentations, Resize9 10 11class Datasets(Dataset):12    def __init__(self, data_file, transform=None, phase='train', *args, **kwargs):13        self.transform = transform14        self.data_info = pd.read_csv(data_file, index_col=0)15        self.phase = phase16 17    def __len__(self):18        return len(self.data_info)19 20    def __getitem__(self, index):21        data = self.pull_item_seg(index)22        return data23 24    def pull_item_seg(self, index):25        """26        :param index: image index27        """28        data = self.data_info.iloc[index]29        img_name = data['img']30        label_name = data['label']31 32        ori_img = io.imread(img_name, as_gray=False)33        ori_label = io.imread(label_name, as_gray=True)34        assert (ori_img is not None and ori_label is not None), f'{img_name} or {label_name} is not valid'35 36        if self.transform is not None:37            img, label = self.transform((ori_img, ori_label))38 39        one_hot_label = np.zeros([2] + list(label.shape), dtype=np.float)40        one_hot_label[0] = label == 041        one_hot_label[1] = label > 042        return_dict = {43            'img': torch.from_numpy(img).permute(2, 0, 1),44            'label': torch.from_numpy(one_hot_label),45            'img_name': os.path.basename(img_name)46        }47        return return_dict48 49 50def get_data_loader(config, test_mode=False):51    if not test_mode:52        train_params = {53            'batch_size': config['BATCH_SIZE'],54            'shuffle': config['IS_SHUFFLE'],55            'drop_last': False,56            'collate_fn': collate_fn,57            'num_workers': config['NUM_WORKERS'],58            'pin_memory': False59        }60        #  data_file, config, transform=None61        train_set = Datasets(62            config['DATASET'],63            Augmentations(64                config['IMG_SIZE'], config['PRIOR_MEAN'], config['PRIOR_STD'], 'train', config['PHASE'], config65            ),66            config['PHASE'],67            config68        )69        patterns = ['train']70    else:71        patterns = []72 73    if config['IS_VAL']:74        val_params = {75            'batch_size': config['VAL_BATCH_SIZE'],76            'shuffle': False,77            'drop_last': False,78            'collate_fn': collate_fn,79            'num_workers': config['NUM_WORKERS'],80            'pin_memory': False81        }82        val_set = Datasets(83            config['VAL_DATASET'],84            Augmentations(85                config['IMG_SIZE'], config['PRIOR_MEAN'], config['PRIOR_STD'], 'val', config['PHASE'], config86            ),87            config['PHASE'],88            config89        )90        patterns += ['val']91 92    if config['IS_TEST']:93        test_params = {94            'batch_size': config['VAL_BATCH_SIZE'],95            'shuffle': False,96            'drop_last': False,97            'collate_fn': collate_fn,98            'num_workers': config['NUM_WORKERS'],99            'pin_memory': False100        }101        test_set = Datasets(102            config['TEST_DATASET'],103            Augmentations(104                config['IMG_SIZE'], config['PRIOR_MEAN'], config['PRIOR_STD'], 'test', config['PHASE'], config105            ),106            config['PHASE'],107            config108        )109        patterns += ['test']110 111    data_loaders = {}112    for x in patterns:113        data_loaders[x] = DataLoader(eval(x+'_set'), **eval(x+'_params'))114    return data_loaders115 116 117def collate_fn(batch):118    def to_tensor(item):119        if torch.is_tensor(item):120            return item121        elif isinstance(item, type(np.array(0))):122            return torch.from_numpy(item).float()123        elif isinstance(item, type('0')):124            return item125        elif isinstance(item, list):126            return item127        elif isinstance(item, dict):128            return item129 130    return_data = {}131    for key in batch[0].keys():132        return_data[key] = []133 134    for sample in batch:135        for key, value in sample.items():136            return_data[key].append(to_tensor(value))137 138    keys = set(batch[0].keys()) - {'img_name'}139    for key in keys:140        return_data[key] = torch.stack(return_data[key], dim=0)141 142    return return_data143 144