| """Shared configuration and device helpers for SEEDS scripts.""" |
|
|
| from __future__ import annotations |
|
|
| import random |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| MODEL_DIR = PROJECT_ROOT / "model" |
| if str(MODEL_DIR) not in sys.path: |
| sys.path.insert(0, str(MODEL_DIR)) |
|
|
| from seeds import SEEDS |
|
|
|
|
| def load_config(path: str) -> dict: |
| with open(path, "r", encoding="utf-8") as handle: |
| return yaml.safe_load(handle) |
|
|
|
|
| def choose_device(value: str) -> torch.device: |
| if value == "auto": |
| return torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| return torch.device(value) |
|
|
|
|
| def set_seed(seed: int) -> None: |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| if torch.cuda.is_available(): |
| torch.cuda.manual_seed_all(seed) |
|
|
|
|
| def build_model(config: dict) -> SEEDS: |
| data, model = config["data"], config["model"] |
| return SEEDS( |
| channels=len(data["variables"]), faces=data["faces"], height=data["height"], width=data["width"], |
| seed_count=data["seed_count"], patch_size=model["patch_size"], embed_dim=model["embed_dim"], |
| spatial_layers=model["spatial_layers"], field_layers=model["field_layers"], |
| sequence_layers=model["sequence_layers"], mlp_ratio=model["mlp_ratio"], dropout=model["dropout"], |
| sigma_min=model["sigma_min"], sigma_max=model["sigma_max"], |
| ) |
|
|
|
|
| def resolve_path(path: str, config_path: str) -> Path: |
| candidate = Path(path) |
| if candidate.is_absolute(): |
| return candidate |
| return Path(config_path).resolve().parent.parent / candidate |
|
|