KyanChen/BuildingExtraction
5
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 