| 266 | pin_memory=True) |
| 267 | |
| 268 | def test_dataloader(self): |
| 269 | dataset = dataset_dict[self.hparams.dataset_name] |
| 270 | kwargs = { |
| 271 | 'root_dir': self.hparams.root_dir, |
| 272 | 'img_wh': tuple(self.hparams.img_wh), |
| 273 | 'mask_dir': self.hparams.mask_dir, |
| 274 | 'canonical_wh': self.hparams.canonical_wh, |
| 275 | 'canonical_dir': self.hparams.canonical_dir, |
| 276 | 'test': self.hparams.test |
| 277 | } |
| 278 | self.train_dataset = dataset(split='train', **kwargs) |
| 279 | return DataLoader( |
| 280 | self.train_dataset, |
| 281 | shuffle=False, |
| 282 | num_workers=4, |
| 283 | batch_size=1, # validate one image (H*W rays) at a time. |
| 284 | pin_memory=True) |
| 285 | |
| 286 | def training_step(self, batch, batch_idx): |
| 287 | # Fetch training data. |