| 22 | |
| 23 | |
| 24 | class DocData(Dataset): |
| 25 | def __init__(self, path_img, path_gt, loadSize, mode=1): |
| 26 | super().__init__() |
| 27 | self.path_gt = path_gt |
| 28 | self.path_img = path_img |
| 29 | self.data_gt = os.listdir(path_gt) |
| 30 | self.data_img = os.listdir(path_img) |
| 31 | self.mode = mode |
| 32 | if mode == 1: |
| 33 | self.ImgTrans = (ImageTransform(loadSize)["train"], ImageTransform(loadSize)["train_gt"]) |
| 34 | else: |
| 35 | self.ImgTrans = ImageTransform(loadSize)["test"] |
| 36 | |
| 37 | def __len__(self): |
| 38 | return len(self.data_gt) |
| 39 | |
| 40 | def __getitem__(self, idx): |
| 41 | |
| 42 | gt = Image.open(os.path.join(self.path_gt, self.data_img[idx])) |
| 43 | img = Image.open(os.path.join(self.path_img, self.data_img[idx])) |
| 44 | img = img.convert('RGB') |
| 45 | gt = gt.convert('RGB') |
| 46 | if self.mode == 1: |
| 47 | seed = torch.random.seed() |
| 48 | torch.random.manual_seed(seed) |
| 49 | img = self.ImgTrans[0](img) |
| 50 | torch.random.manual_seed(seed) |
| 51 | gt = self.ImgTrans[1](gt) |
| 52 | else: |
| 53 | img= self.ImgTrans(img) |
| 54 | gt = self.ImgTrans(gt) |
| 55 | name = self.data_img[idx] |
| 56 | return img, gt, name |