File size: 4,312 Bytes
590a501 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | """Centralized project settings loaded from YAML + environment overrides."""
from __future__ import annotations
import os
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import yaml
PROJECT_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_BASE_CONFIG = PROJECT_ROOT / "config" / "base.yaml"
def _deep_get(d: dict, *keys, default=None):
cur = d
for k in keys:
if not isinstance(cur, dict) or k not in cur:
return default
cur = cur[k]
return cur
@dataclass
class ProjectSettings:
raw: dict[str, Any]
project_root: Path = PROJECT_ROOT
@property
def qlib_provider_uri(self) -> str:
uri = self.raw["qlib"]["provider_uri"]
p = Path(uri)
if not p.is_absolute():
p = self.project_root / p
return str(p)
@property
def qlib_region(self) -> str:
return self.raw["qlib"].get("region", "cn")
@property
def market(self) -> str:
return self.raw["qlib"].get("market", "csi300")
@property
def freq(self) -> str:
return self.raw.get("data", {}).get("freq", "day")
@property
def backtest_freq(self) -> str:
bt_freq = self.raw.get("backtest", {}).get("freq")
if bt_freq:
return bt_freq
return self.freq
@property
def benchmark(self) -> str | None:
bench = self.raw.get("backtest", {}).get("benchmark", "SH000300")
if bench in (None, "null", "none", ""):
return None
return bench
@property
def segments(self) -> dict[str, tuple[str, str]]:
splits = self.raw.get("splits", {})
return {k: (v["start"], v["end"]) for k, v in splits.items()}
@property
def fit_segment(self) -> tuple[str, str]:
return self.segments["train"]
@property
def output_root(self) -> Path:
root = Path(_deep_get(self.raw, "output", "root", default="outputs"))
if not root.is_absolute():
root = self.project_root / root
return root
@property
def mlruns_uri(self) -> str:
uri = _deep_get(self.raw, "experiment", "mlruns_uri", default="mlruns")
p = Path(uri)
if not p.is_absolute():
p = self.project_root / p
return f"file://{p}"
@property
def dump_config(self) -> dict[str, Any]:
return self.raw.get("dump", {})
@property
def backtest_config(self) -> dict[str, Any]:
return self.raw.get("backtest", {})
def path(self, *parts: str) -> Path:
return self.project_root.joinpath(*parts)
def gp_output_dir(self, run_id: str | None = None) -> Path:
run_id = run_id or os.environ.get("RUN_ID", "qlib_gp_run_0")
return self.output_root / "gp_mining" / run_id
def gp_dataset_path(self, run_id: str | None = None) -> Path:
return self.gp_output_dir(run_id) / "gp_qlib_dataset.pkl"
def workflow_path(self, name: str) -> Path:
return self.project_root / "config" / "workflows" / name
@property
def strategy_config_path(self) -> Path:
return self.project_root / "config" / "strategies.yaml"
@property
def quantaalpha_config_path(self) -> Path:
return self.project_root / "config" / "quantaalpha.yaml"
@property
def factor_registry_path(self) -> Path:
return self.project_root / "config" / "factor_registry.yaml"
def factor_registry_output_dir(self) -> Path:
return self.output_root / "factors" / "registry"
def backtest_output_dir(self, name: str = "default") -> Path:
return self.output_root / "backtest" / name
def load_settings(config_path: str | Path | None = None) -> ProjectSettings:
path = Path(config_path) if config_path else DEFAULT_BASE_CONFIG
with open(path, encoding="utf-8") as f:
raw = yaml.safe_load(f)
# Environment overrides (avoid hardcoding in scripts)
if "QLIB_PROVIDER_URI" in os.environ:
raw.setdefault("qlib", {})["provider_uri"] = os.environ["QLIB_PROVIDER_URI"]
if "QLIB_MARKET" in os.environ:
raw.setdefault("qlib", {})["market"] = os.environ["QLIB_MARKET"]
if "RUN_ID" in os.environ:
raw.setdefault("experiment", {})["run_id"] = os.environ["RUN_ID"]
return ProjectSettings(raw=raw)
|