SACFlow / sacflow /utils /config.py
sathiiii's picture
Add SACFlow source code
d65ae7d verified
Raw
History Blame Contribute Delete
1.56 kB
from __future__ import annotations
from pathlib import Path
from typing import Any, Dict, List
import copy
import yaml
def _deep_update(base: Dict[str, Any], update: Dict[str, Any]) -> Dict[str, Any]:
out = copy.deepcopy(base)
for k, v in update.items():
if isinstance(v, dict) and isinstance(out.get(k), dict):
out[k] = _deep_update(out[k], v)
else:
out[k] = copy.deepcopy(v)
return out
def load_yaml(path: str | Path) -> Dict[str, Any]:
path = Path(path)
with open(path, "r") as f:
cfg = yaml.safe_load(f) or {}
base_files = cfg.pop("_base_", [])
if isinstance(base_files, str):
base_files = [base_files]
merged: Dict[str, Any] = {}
for base in base_files:
base_path = (path.parent / base).resolve()
merged = _deep_update(merged, load_yaml(base_path))
merged = _deep_update(merged, cfg)
return merged
def save_yaml(cfg: Dict[str, Any], path: str | Path) -> None:
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w") as f:
yaml.safe_dump(cfg, f, sort_keys=False)
def get(cfg: Dict[str, Any], key: str, default=None):
cur = cfg
for part in key.split("."):
if not isinstance(cur, dict) or part not in cur:
return default
cur = cur[part]
return cur
def set_by_path(cfg: Dict[str, Any], key: str, value: Any) -> None:
cur = cfg
parts = key.split(".")
for p in parts[:-1]:
cur = cur.setdefault(p, {})
cur[parts[-1]] = value