Download scripts/common.py from OneScience-Group/FourCastNet_v2: direct link, hf CLI and curl.
- Browser
- Download file 2.99 kB
-
https://huggingface.co/OneScience-Group/FourCastNet_v2/resolve/main/scripts/common.py
- Command line
-
hf download hf://OneScience-Group/FourCastNet_v2/scripts/common.py
-
curl -L -o common.py https://huggingface.co/OneScience-Group/FourCastNet_v2/resolve/main/scripts/common.py
2.99 kB
| from __future__ import annotations | |
| import random | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import torch | |
| import yaml | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| DEFAULT_CONFIG = PROJECT_ROOT / "conf" / "config.yaml" | |
| def load_config(path: str | Path = DEFAULT_CONFIG) -> dict[str, Any]: | |
| config_path = Path(path).expanduser().resolve() | |
| with config_path.open("r", encoding="utf-8") as stream: | |
| config = yaml.safe_load(stream) | |
| config["_config_path"] = str(config_path) | |
| config["_project_root"] = str(PROJECT_ROOT) | |
| validate_config(config) | |
| return config | |
| def validate_config(config: dict[str, Any]) -> None: | |
| variables = config["data"]["variables"] | |
| if len(variables) != 73 or len(set(variables)) != 73: | |
| raise ValueError("FourCastNet v2 requires 73 unique variables") | |
| profile_name = config["model"]["profile"] | |
| profiles = config["model"]["profiles"] | |
| if profile_name not in profiles: | |
| raise ValueError(f"Unknown model profile: {profile_name}") | |
| profile = profiles[profile_name] | |
| if profile["in_channels"] != len(variables): | |
| raise ValueError("Model input channels do not match the variable ledger") | |
| if profile["out_channels"] != len(variables): | |
| raise ValueError("Model output channels do not match the variable ledger") | |
| if config["data"]["input_steps"] != 1: | |
| raise ValueError("FourCastNet v2 expects exactly one input time step") | |
| if config["data"]["output_steps"] != 1: | |
| raise ValueError("One-step pretraining expects data.output_steps=1") | |
| if config["training"]["finetune"]["autoregressive_steps"] < 2: | |
| raise ValueError("Fine-tuning requires at least two autoregressive steps") | |
| if config["inference"]["rollout_steps"] < 1: | |
| raise ValueError("inference.rollout_steps must be positive") | |
| if config["training"]["stage"] not in {"one_step", "finetune"}: | |
| raise ValueError("training.stage must be 'one_step' or 'finetune'") | |
| if config["checkpoint"]["initialize_from"] != "scratch": | |
| raise ValueError("checkpoint.initialize_from must be 'scratch'") | |
| if not config["checkpoint"].get("finetune_from"): | |
| raise ValueError("checkpoint.finetune_from must name a one-step checkpoint") | |
| prefix = config["checkpoint"].get("prefix", "model_bak") | |
| if not prefix or Path(prefix).name != prefix: | |
| raise ValueError("checkpoint.prefix must be a non-empty file name") | |
| def resolve_path(config: dict[str, Any], value: str | Path) -> Path: | |
| path = Path(value).expanduser() | |
| if path.is_absolute(): | |
| return path | |
| return Path(config["_project_root"]) / path | |
| def active_model_config(config: dict[str, Any]) -> dict[str, Any]: | |
| return dict(config["model"]["profiles"][config["model"]["profile"]]) | |
| def seed_everything(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) | |