Download scripts/common.py from OneScience-Group/SEEDS: direct link, hf CLI and curl.
- Browser
- Download file 1.67 kB
-
https://huggingface.co/OneScience-Group/SEEDS/resolve/main/scripts/common.py
- Command line
-
hf download hf://OneScience-Group/SEEDS/scripts/common.py
-
curl -L -o common.py https://huggingface.co/OneScience-Group/SEEDS/resolve/main/scripts/common.py
1.67 kB
| """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 | |