Download utils/data.py from gaaaaaaaaaaa/multimodal-reasoning: direct link, hf CLI and curl.
- Browser
- Download file 1.55 kB
-
https://huggingface.co/gaaaaaaaaaaa/multimodal-reasoning/resolve/main/utils/data.py
- Command line
-
hf download hf://gaaaaaaaaaaa/multimodal-reasoning/utils/data.py
-
curl -L -o data.py https://huggingface.co/gaaaaaaaaaaa/multimodal-reasoning/resolve/main/utils/data.py
1.55 kB
| from typing import Callable, Any | |
| from datasets import load_dataset | |
| import utils.configs_loader as ProjectConfigs | |
| # PUBLIC CLASS | |
| class StreamingDataset: | |
| def __init__( | |
| self, | |
| path: str, | |
| split: str = "train", | |
| mapping: Callable[[dict], dict] | None = None, | |
| shuffle: bool = True, | |
| buffer_size: int = 10_000, | |
| seed: int = 42, | |
| ): | |
| self.dataset_info = ProjectConfigs.get(("DATASET", path)) | |
| self.split = split | |
| self.mapping = mapping | |
| self.should_shuffle = shuffle | |
| self.buffer_size = buffer_size | |
| self.seed = seed | |
| self.dataset = None | |
| def take(self, amount: int): | |
| if self.dataset is None: | |
| self.load() | |
| return self.dataset.take(amount) | |
| def shuffle(self): | |
| if self.dataset is None: | |
| raise RuntimeError("Dataset must be loaded before shuffle().") | |
| self.dataset = self.dataset.shuffle( | |
| seed=self.seed, | |
| buffer_size=self.buffer_size, | |
| ) | |
| return self.dataset | |
| def set_map(self, method: Callable[[dict], dict] | None): | |
| self.mapping = method | |
| def load(self): | |
| self.dataset = load_dataset( | |
| path=self.dataset_info.Path, | |
| name=self.dataset_info.Name, | |
| split=self.split, | |
| streaming=True, | |
| ) | |
| if self.should_shuffle: | |
| self.shuffle() | |
| if self.mapping is not None: | |
| self.dataset = self.dataset.map(self.mapping) | |
| return self.dataset |