|
1 | 1 | from __future__ import annotations |
2 | 2 |
|
| 3 | +from dataclasses import dataclass |
3 | 4 | from pathlib import Path |
4 | | -from typing import TYPE_CHECKING |
| 5 | +from typing import Any |
5 | 6 |
|
6 | 7 | import torch |
7 | 8 | from lightning.pytorch import LightningDataModule |
8 | | -from torch.utils.data import DataLoader, Dataset |
| 9 | +from torch.utils.data import ConcatDataset, DataLoader, Dataset |
9 | 10 |
|
10 | 11 | from electrai.dataloader import utils |
11 | 12 | from electrai.dataloader.collate import collate_fn |
12 | 13 | from electrai.dataloader.split import split_data |
13 | 14 |
|
14 | | -if TYPE_CHECKING: |
15 | | - import os |
16 | | - |
17 | 15 | dtype_map = {"f32": torch.float32, "f16": torch.float16, "bf16": torch.bfloat16} |
18 | 16 |
|
19 | 17 |
|
| 18 | +@dataclass(frozen=True) |
| 19 | +class DatasetSpec: |
| 20 | + root: str |
| 21 | + split_file: str | None = None |
| 22 | + val_frac: float | None = None |
| 23 | + dataset_id: int | None = None |
| 24 | + |
| 25 | + |
| 26 | +class AddDatasetID(Dataset): |
| 27 | + """Wrap dataset so every sample carries a constant Dataset_ID""" |
| 28 | + |
| 29 | + def __init__(self, base: Dataset, functional_id: int, key: str = "Dataset_ID"): |
| 30 | + self.base = base |
| 31 | + self.functional_id = int(functional_id) |
| 32 | + self.key = key |
| 33 | + |
| 34 | + def __len__(self): |
| 35 | + return len(self.base) |
| 36 | + |
| 37 | + def __getitem__(self, idx): |
| 38 | + out = self.base[idx] |
| 39 | + out[self.key] = self.functional_id |
| 40 | + return out |
| 41 | + |
| 42 | + |
20 | 43 | class RhoRead(LightningDataModule): |
| 44 | + """ |
| 45 | + Works based on these keys for one or more datasets: |
| 46 | + `root`, `split_file`, `val_frac` (if `split_file` is null) and dataset_id |
| 47 | + """ |
| 48 | + |
21 | 49 | def __init__( |
22 | 50 | self, |
23 | | - root: str | bytes | os.PathLike, |
24 | | - precision: str, |
| 51 | + datasets: list[dict[str, Any]] | list[DatasetSpec] | None = None, |
| 52 | + default_val_frac: float = 0.005, |
| 53 | + default_split_file: str | None = None, |
| 54 | + precision: str = "f32", |
25 | 55 | batch_size: int = 2, |
26 | 56 | train_workers: int = 8, |
27 | 57 | val_workers: int = 2, |
28 | 58 | pin_memory: bool = False, |
29 | | - val_frac: float = 0.005, |
30 | 59 | drop_last: bool = False, |
31 | | - split_file: str | bytes | os.PathLike | None = None, |
32 | 60 | augmentation: bool = False, |
33 | 61 | random_seed: int = 42, |
34 | 62 | ): |
35 | 63 | super().__init__() |
36 | 64 | self.save_hyperparameters() |
37 | | - self.root = root |
38 | 65 | self.batch_size = batch_size |
39 | 66 | self.train_workers = train_workers |
40 | 67 | self.val_workers = val_workers |
41 | 68 | self.pin_memory = pin_memory |
42 | | - self.val_frac = val_frac |
43 | 69 | self.drop_last = drop_last |
44 | | - self.split_file = split_file |
45 | 70 | self.precision = precision |
46 | 71 | self.augmentation = augmentation |
47 | 72 | self.random_seed = random_seed |
48 | 73 |
|
49 | | - def setup(self, stage=None): |
50 | | - dataset = RhoData( |
51 | | - self.root, precision=self.precision, augmentation=self.augmentation |
52 | | - ) |
53 | | - self.subsets = split_data( |
54 | | - dataset, |
55 | | - val_frac=self.val_frac, |
56 | | - split_file=self.split_file, |
57 | | - random_seed=self.random_seed, |
58 | | - ) |
| 74 | + if datasets is None or len(datasets) == 0: |
| 75 | + raise ValueError("`datasets` must contain at least one dataset spec.") |
| 76 | + |
| 77 | + specs: list[DatasetSpec] = [ |
| 78 | + d if isinstance(d, DatasetSpec) else DatasetSpec(**d) for d in datasets |
| 79 | + ] |
| 80 | + filled: list[DatasetSpec] = [] |
| 81 | + for i, s in enumerate(specs): |
| 82 | + if s.root is None: |
| 83 | + raise ValueError( |
| 84 | + f"`root` is required for dataset {i + 1}. Received: {s.root!r}" |
| 85 | + ) |
| 86 | + dataset_id = s.dataset_id if s.dataset_id is not None else i |
| 87 | + filled.append( |
| 88 | + DatasetSpec( |
| 89 | + root=s.root, |
| 90 | + split_file=( |
| 91 | + s.split_file if s.split_file is not None else default_split_file |
| 92 | + ), |
| 93 | + val_frac=( |
| 94 | + s.val_frac if s.val_frac is not None else default_val_frac |
| 95 | + ), |
| 96 | + dataset_id=dataset_id, |
| 97 | + ) |
| 98 | + ) |
| 99 | + self.specs = filled |
| 100 | + |
| 101 | + self.train_set: Dataset | None = None |
| 102 | + self.val_set: Dataset | None = None |
| 103 | + self.test_set: Dataset | None = None |
| 104 | + |
| 105 | + def setup(self, stage: str): |
| 106 | + train_parts: list[Dataset] = [] |
| 107 | + val_parts: list[Dataset] = [] |
| 108 | + test_parts: list[Dataset] = [] |
| 109 | + |
| 110 | + for spec in self.specs: |
| 111 | + ds = RhoData( |
| 112 | + spec.root, precision=self.precision, augmentation=self.augmentation |
| 113 | + ) |
| 114 | + splits = split_data( |
| 115 | + ds, |
| 116 | + val_frac=float(spec.val_frac), |
| 117 | + split_file=spec.split_file, |
| 118 | + random_seed=self.random_seed, |
| 119 | + ) |
| 120 | + |
| 121 | + dataset_id = int(spec.dataset_id) |
| 122 | + |
| 123 | + train_parts.append(AddDatasetID(splits["train"], dataset_id)) |
| 124 | + val_parts.append(AddDatasetID(splits["validation"], dataset_id)) |
| 125 | + if "test" in splits and splits["test"] is not None: |
| 126 | + test_parts.append(AddDatasetID(splits["test"], dataset_id)) |
| 127 | + |
59 | 128 | if stage == "fit": |
60 | | - self.train_set = self.subsets["train"] |
61 | | - self.val_set = self.subsets["validation"] |
| 129 | + self.train_set = ( |
| 130 | + train_parts[0] if len(train_parts) == 1 else ConcatDataset(train_parts) |
| 131 | + ) |
| 132 | + self.val_set = ( |
| 133 | + val_parts[0] if len(val_parts) == 1 else ConcatDataset(val_parts) |
| 134 | + ) |
62 | 135 | elif stage == "test": |
63 | | - self.test_set = self.subsets["test"] |
| 136 | + self.test_set = ( |
| 137 | + test_parts[0] if len(test_parts) == 1 else ConcatDataset(test_parts) |
| 138 | + ) |
64 | 139 |
|
65 | 140 | def train_dataloader(self): |
66 | 141 | return DataLoader( |
|
0 commit comments