(train_cfg: Config)
| 11 | |
| 12 | |
| 13 | def test_create_dataloader_cache(train_cfg: Config): |
| 14 | train_cfg.task.data.shuffle = False |
| 15 | train_cfg.task.data.batch_size = 2 |
| 16 | |
| 17 | cache_file = Path("tests/data/train.cache") |
| 18 | cache_file.unlink(missing_ok=True) |
| 19 | |
| 20 | make_cache_loader = create_dataloader(train_cfg.task.data, train_cfg.dataset) |
| 21 | load_cache_loader = create_dataloader(train_cfg.task.data, train_cfg.dataset) |
| 22 | m_batch_size, m_images, _, m_reverse_tensors, m_image_paths = next(iter(make_cache_loader)) |
| 23 | l_batch_size, l_images, _, l_reverse_tensors, l_image_paths = next(iter(load_cache_loader)) |
| 24 | assert m_batch_size == l_batch_size |
| 25 | assert m_images.shape == l_images.shape |
| 26 | assert m_reverse_tensors.shape == l_reverse_tensors.shape |
| 27 | assert m_image_paths == l_image_paths |
| 28 | |
| 29 | |
| 30 | def test_training_data_loader_correctness(train_dataloader: DataLoader): |
nothing calls this directly
no test coverage detected