CLYang617/RemoteSensingChangeDetection-RSCD.HA2F
0
1import cv22import os3from os.path import join as osp4import numpy5import torch.utils.data6 7 8class Dataset(torch.utils.data.Dataset):9 def __init__(self, file_root='data/', mode='train', transform=None):10 self.file_list = os.listdir(osp(file_root, mode, 'A'))11 12 self.pre_images = [osp(file_root, mode, 'A', x) for x in self.file_list]13 self.post_images = [osp(file_root, mode, 'B', x) for x in self.file_list]14 self.gts = [osp(file_root, mode, 'label', x) for x in self.file_list]15 16 self.transform = transform17 18 def __len__(self):19 return len(self.pre_images)20 21 def __getitem__(self, idx):22 pre_image_name = self.pre_images[idx]23 label_name = self.gts[idx]24 post_image_name = self.post_images[idx]25 26 pre_image = cv2.imread(pre_image_name)27 label = cv2.imread(label_name, 0)28 post_image = cv2.imread(post_image_name)29 30 img = numpy.concatenate((pre_image, post_image), axis=2)31 32 if self.transform:33 [img, label] = self.transform(img, label)34 35 return img, label36 37 def get_img_info(self, idx):38 img = cv2.imread(self.pre_images[idx])39 return {"height": img.shape[0], "width": img.shape[1]}40 