Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions scripts/e2e_train.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,12 +186,11 @@ class Config:
raise ValueError("no samples remain after filtering")

datamodule = RhoRead(
root=str(filelist),
datasets=[{"root": str(filelist), "val_frac": 0.4}],
precision="f32",
batch_size=cfg.nbatch,
train_workers=0,
val_workers=0,
val_frac=0.4,
augmentation=False,
random_seed=seed,
)
Expand Down
9 changes: 4 additions & 5 deletions src/electrai/configs/MP/config_resnet.yaml
Original file line number Diff line number Diff line change
@@ -1,19 +1,18 @@
# Dataset / loader parameters
data:
_target_: electrai.dataloader.dataset.RhoRead
root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/mp_filelist.txt
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/split_limit_22M.json
datasets:
- root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/mp_filelist.txt
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/split.json
val_frac: 0.005 # ignored when split_file is provided
precision: f32
batch_size: 1
train_workers: 8
val_workers: 2
pin_memory: false
val_frac: 0.005
drop_last: false
augmentation: false
random_seed: 42
# downsample_label: 0
# downsample_data: 0

# Model
model:
Expand Down
9 changes: 4 additions & 5 deletions src/electrai/configs/MP/config_resunet.yaml
Original file line number Diff line number Diff line change
@@ -1,19 +1,18 @@
# Dataset / loader parameters
data:
_target_: electrai.dataloader.dataset.RhoRead
root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/mp_filelist.txt
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/rho_gga/split_limit_22M.json
datasets:
- root: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/mp_filelist.txt
split_file: /scratch/gpfs/ROSENGROUP/common/globus_share_OA/mp/chg_datasets/dataset_2/split.json
val_frac: 0.005 # ignored when split_file is provided
precision: f32
batch_size: 1
train_workers: 8
val_workers: 2
pin_memory: false
val_frac: 0.005
drop_last: false
augmentation: false
random_seed: 42
# downsample_label: 0
# downsample_data: 0

# Model
model:
Expand Down
9 changes: 7 additions & 2 deletions src/electrai/dataloader/collate.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,5 +7,10 @@ def collate_fn(batch):
try:
return default_collate(batch)
except RuntimeError:
x, y, index = zip(*batch, strict=True)
return list(x), list(y), list(index)
result = {}
for k in batch[0]:
try:
result[k] = default_collate([d[k] for d in batch])
except RuntimeError:
result[k] = [d[k] for d in batch]
return result
138 changes: 113 additions & 25 deletions src/electrai/dataloader/dataset.py
Original file line number Diff line number Diff line change
@@ -1,66 +1,154 @@
from __future__ import annotations

import warnings
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING
from typing import Any

import torch
from lightning.pytorch import LightningDataModule
from torch.utils.data import DataLoader, Dataset
from torch.utils.data import ConcatDataset, DataLoader, Dataset

from electrai.dataloader import utils
from electrai.dataloader.collate import collate_fn
from electrai.dataloader.split import split_data

if TYPE_CHECKING:
import os

dtype_map = {"f32": torch.float32, "f16": torch.float16, "bf16": torch.bfloat16}


@dataclass(frozen=True)
class DatasetSpec:
root: str
split_file: str | None = None
val_frac: float | None = None
dataset_id: int | None = None


class AddDatasetID(Dataset):
"""Wrap dataset so every sample carries a constant Dataset_ID"""

def __init__(self, base: Dataset, functional_id: int, key: str = "Dataset_ID"):
self.base = base
self.functional_id = int(functional_id)
self.key = key

def __len__(self):
return len(self.base)

def __getitem__(self, idx):
out = dict(self.base[idx])
out[self.key] = self.functional_id
return out


class RhoRead(LightningDataModule):
"""
Works based on these keys for one or more datasets:
`root`, `split_file`, `val_frac` (if `split_file` is null) and dataset_id
"""

def __init__(
self,
root: str | bytes | os.PathLike,
precision: str,
datasets: list[dict[str, Any]] | list[DatasetSpec] | None = None,
default_val_frac: float = 0.005,
default_split_file: str | None = None,
precision: str = "f32",
batch_size: int = 2,
train_workers: int = 8,
val_workers: int = 2,
pin_memory: bool = False,
val_frac: float = 0.005,
drop_last: bool = False,
split_file: str | bytes | os.PathLike | None = None,
augmentation: bool = False,
random_seed: int = 42,
):
super().__init__()
self.save_hyperparameters()
self.root = root
self.batch_size = batch_size
self.train_workers = train_workers
self.val_workers = val_workers
self.pin_memory = pin_memory
self.val_frac = val_frac
self.drop_last = drop_last
self.split_file = split_file
self.precision = precision
self.augmentation = augmentation
self.random_seed = random_seed

def setup(self, stage=None):
dataset = RhoData(
self.root, precision=self.precision, augmentation=self.augmentation
)
self.subsets = split_data(
dataset,
val_frac=self.val_frac,
split_file=self.split_file,
random_seed=self.random_seed,
)
if datasets is None or len(datasets) == 0:
raise ValueError("`datasets` must contain at least one dataset spec.")

specs: list[DatasetSpec] = [
d if isinstance(d, DatasetSpec) else DatasetSpec(**d) for d in datasets
]
filled: list[DatasetSpec] = []
for i, s in enumerate(specs):
dataset_id = s.dataset_id if s.dataset_id is not None else i
filled.append(
DatasetSpec(
root=s.root,
split_file=(
s.split_file if s.split_file is not None else default_split_file
),
val_frac=(
s.val_frac if s.val_frac is not None else default_val_frac
),
dataset_id=dataset_id,
)
)
self.specs = filled

self.train_set: Dataset | None = None
self.val_set: Dataset | None = None
self.test_set: Dataset | None = None

def setup(self, stage: str):
train_parts: list[Dataset] = []
val_parts: list[Dataset] = []
test_parts: list[Dataset] = []

for spec in self.specs:
if stage == "test" and spec.split_file is None:
warnings.warn(
f"Dataset with root '{spec.root}' has no split_file and will be "
"skipped for the test stage (test indices require a split_file).",
UserWarning,
stacklevel=2,
)
continue

ds = RhoData(
spec.root, precision=self.precision, augmentation=self.augmentation
)
dataset_id = int(spec.dataset_id)

if stage == "fit":
splits = split_data(
ds,
val_frac=float(spec.val_frac),
split_file=spec.split_file,
random_seed=self.random_seed,
)
train_parts.append(AddDatasetID(splits["train"], dataset_id))
val_parts.append(AddDatasetID(splits["validation"], dataset_id))
elif stage == "test":
splits = split_data(
ds,
val_frac=float(spec.val_frac),
split_file=spec.split_file,
random_seed=self.random_seed,
)
if "test" in splits and splits["test"] is not None:
test_parts.append(AddDatasetID(splits["test"], dataset_id))

if stage == "fit":
self.train_set = self.subsets["train"]
self.val_set = self.subsets["validation"]
self.train_set = (
train_parts[0] if len(train_parts) == 1 else ConcatDataset(train_parts)
)
self.val_set = (
val_parts[0] if len(val_parts) == 1 else ConcatDataset(val_parts)
)
elif stage == "test":
self.test_set = self.subsets["test"]
self.test_set = (
test_parts[0] if len(test_parts) == 1 else ConcatDataset(test_parts)
)

def train_dataloader(self):
return DataLoader(
Expand Down
11 changes: 9 additions & 2 deletions src/electrai/entrypoints/test.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import os
from pathlib import Path
from types import SimpleNamespace

Expand Down Expand Up @@ -48,12 +49,18 @@ def test(args):
tmp_dir = log_dir / "tmp"
for directory in [log_dir, tmp_dir]:
directory.mkdir(exist_ok=True, parents=True)
local_world_size = int(
os.environ.get("LOCAL_WORLD_SIZE", torch.cuda.device_count())
)
world_size = int(os.environ.get("WORLD_SIZE", local_world_size))
num_nodes = max(1, world_size // local_world_size)
trainer = Trainer(
logger=None,
callbacks=None,
accelerator="gpu" if torch.cuda.is_available() else "cpu",
devices=1,
precision=cfg.model_precision,
devices="auto",
num_nodes=num_nodes,
precision=cfg.precision,
)

lit_model.test_cfg = SimpleNamespace(
Expand Down
21 changes: 18 additions & 3 deletions src/electrai/lightning.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ def validation_step(self, batch):
return loss

def _loss_calculation(self, batch):
# batch["Dataset_ID"] is available for future multi-head model extensions
x = batch["data"]
y = batch["label"]
if isinstance(x, list):
Expand Down Expand Up @@ -94,17 +95,31 @@ def test_step(self, batch):
end = torch.cuda.Event(enable_timing=True)

start.record()
preds = self(x)
loss = self.loss_fn(preds, y)
if isinstance(x, list):
preds_list, losses = [], []
for x_i, y_i in zip(x, y, strict=True):
p = self(x_i.unsqueeze(0))
preds_list.append(p)
losses.append(self.loss_fn(p, y_i.unsqueeze(0)))
preds = torch.cat(preds_list)
loss = torch.stack(losses).mean()
else:
preds = self(x)
loss = self.loss_fn(preds, y)
end.record()

torch.cuda.synchronize()
elapsed = start.elapsed_time(end)

self.log("test_loss", loss, prog_bar=True, sync_dist=True)

y_cpu = (
torch.cat([t.unsqueeze(0) for t in y]).detach().cpu()
if isinstance(y, list)
else y.detach().cpu()
)
out = {
"target": y.detach().cpu(),
"target": y_cpu,
"index": indices,
"nmae": loss.detach().cpu(),
"duration": elapsed,
Expand Down
Loading
Loading