File size: 1,340 Bytes
4947683
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
"""Copyright (c) Meta Platforms, Inc. and affiliates."""

from __future__ import annotations

import os
from pathlib import Path
from typing import Literal

from hydra import compose, initialize_config_dir
from omegaconf import DictConfig
from torch_geometric.loader import DataLoader

from flowmm.model.eval_utils import get_loaders

dataset_options = Literal["carbon", "mp_20", "mpts_52", "perov", "mp_20_llama"]


def init_cfg(
    overrides: list[str] = [],
) -> DictConfig:
    project_root = config_dir = Path(__file__).parents[2]
    os.environ["PROJECT_ROOT"] = str(project_root.resolve())
    config_dir = project_root / f"scripts_model/conf"

    with initialize_config_dir(str(config_dir.resolve()), version_base="1.1"):
        cfg = compose(config_name="default", overrides=overrides)
    return cfg


def init_loaders(
    dataset: dataset_options,
    batch_size: int | None = None,
) -> tuple[DataLoader, DataLoader, DataLoader]:
    overrides = [f"data={dataset}"]
    if batch_size is not None:
        overrides.extend(
            [
                f"data.datamodule.batch_size.train={batch_size}",
                f"data.datamodule.batch_size.val={batch_size}",
                f"data.datamodule.batch_size.test={batch_size}",
            ]
        )
    cfg = init_cfg(overrides=overrides)
    return get_loaders(cfg)