FourCastNet_v2 / scripts /common.py
yzt15806542928's picture
Upload folder using huggingface_hub
eca4864 verified
Raw
History Blame Contribute Delete
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)