Skip to content

Commit c9dd6e3

Browse files
author
Hananeh Oliaei
committed
Addressed handling multiple datasets with the dataloader
1 parent 90590bb commit c9dd6e3

4 files changed

Lines changed: 110 additions & 37 deletions

File tree

src/electrai/configs/MP/config_resnet.yaml

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,19 +1,18 @@
11
# Dataset / loader parameters
22
data:
33
_target_: electrai.dataloader.dataset.RhoRead
4-
root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/mp_filelist.txt
5-
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/split_limit_22M.json
4+
datasets:
5+
- root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/mp_filelist.txt
6+
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/split.json
7+
val_frac: 0.005
68
precision: f32
79
batch_size: 1
810
train_workers: 8
911
val_workers: 2
1012
pin_memory: false
11-
val_frac: 0.005
1213
drop_last: false
1314
augmentation: false
1415
random_seed: 42
15-
# downsample_label: 0
16-
# downsample_data: 0
1716

1817
# Model
1918
model:

src/electrai/configs/MP/config_resunet.yaml

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,19 +1,18 @@
11
# Dataset / loader parameters
22
data:
33
_target_: electrai.dataloader.dataset.RhoRead
4-
root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/mp_filelist.txt
5-
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/split_limit_22M.json
4+
datasets:
5+
- root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/mp_filelist.txt
6+
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/split.json
7+
val_frac: 0.005
68
precision: f32
79
batch_size: 1
810
train_workers: 8
911
val_workers: 2
1012
pin_memory: false
11-
val_frac: 0.005
1213
drop_last: false
1314
augmentation: false
1415
random_seed: 42
15-
# downsample_label: 0
16-
# downsample_data: 0
1716

1817
# Model
1918
model:

src/electrai/dataloader/collate.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,5 +7,5 @@ def collate_fn(batch):
77
try:
88
return default_collate(batch)
99
except RuntimeError:
10-
x, y, index = zip(*batch, strict=True)
11-
return list(x), list(y), list(index)
10+
x, y, index, dataset_id = zip(*batch, strict=True)
11+
return list(x), list(y), list(index), list(dataset_id)

src/electrai/dataloader/dataset.py

Lines changed: 100 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -1,66 +1,141 @@
11
from __future__ import annotations
22

3+
from dataclasses import dataclass
34
from pathlib import Path
4-
from typing import TYPE_CHECKING
5+
from typing import Any
56

67
import torch
78
from lightning.pytorch import LightningDataModule
8-
from torch.utils.data import DataLoader, Dataset
9+
from torch.utils.data import ConcatDataset, DataLoader, Dataset
910

1011
from electrai.dataloader import utils
1112
from electrai.dataloader.collate import collate_fn
1213
from electrai.dataloader.split import split_data
1314

14-
if TYPE_CHECKING:
15-
import os
16-
1715
dtype_map = {"f32": torch.float32, "f16": torch.float16, "bf16": torch.bfloat16}
1816

1917

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+
2043
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+
2149
def __init__(
2250
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",
2555
batch_size: int = 2,
2656
train_workers: int = 8,
2757
val_workers: int = 2,
2858
pin_memory: bool = False,
29-
val_frac: float = 0.005,
3059
drop_last: bool = False,
31-
split_file: str | bytes | os.PathLike | None = None,
3260
augmentation: bool = False,
3361
random_seed: int = 42,
3462
):
3563
super().__init__()
3664
self.save_hyperparameters()
37-
self.root = root
3865
self.batch_size = batch_size
3966
self.train_workers = train_workers
4067
self.val_workers = val_workers
4168
self.pin_memory = pin_memory
42-
self.val_frac = val_frac
4369
self.drop_last = drop_last
44-
self.split_file = split_file
4570
self.precision = precision
4671
self.augmentation = augmentation
4772
self.random_seed = random_seed
4873

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+
59128
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+
)
62135
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+
)
64139

65140
def train_dataloader(self):
66141
return DataLoader(

0 commit comments

Comments
 (0)