Spaces:
Running on Zero
Running on Zero
Download src/dataset/validation_wrapper.py from timfromhcs/AnySplat: direct link, hf CLI and curl.
- Browser
- Download file 1.12 kB
-
https://huggingface.co/spaces/timfromhcs/AnySplat/resolve/main/src/dataset/validation_wrapper.py
- Command line
-
hf download hf://spaces/timfromhcs/AnySplat/src/dataset/validation_wrapper.py
-
curl -L -o validation_wrapper.py https://huggingface.co/spaces/timfromhcs/AnySplat/resolve/main/src/dataset/validation_wrapper.py
1.12 kB
| from typing import Iterator, Optional | |
| import torch | |
| from torch.utils.data import Dataset, IterableDataset | |
| class ValidationWrapper(Dataset): | |
| """Wraps a dataset so that PyTorch Lightning's validation step can be turned into a | |
| visualization step. | |
| """ | |
| dataset: Dataset | |
| dataset_iterator: Optional[Iterator] | |
| length: int | |
| def __init__(self, dataset: Dataset, length: int) -> None: | |
| super().__init__() | |
| self.dataset = dataset | |
| self.length = length | |
| self.dataset_iterator = None | |
| def __len__(self): | |
| return self.length | |
| def __getitem__(self, index: tuple): | |
| if isinstance(self.dataset, IterableDataset): | |
| if self.dataset_iterator is None: | |
| self.dataset_iterator = iter(self.dataset) | |
| return next(self.dataset_iterator) | |
| random_index = torch.randint(0, len(self.dataset), tuple()) | |
| random_context_num = torch.randint(2, self.dataset.view_sampler.num_context_views + 1, tuple()) | |
| # breakpoint() | |
| return self.dataset[random_index.item(), random_context_num.item()] | |