| """Initialize qlib and load project configuration.""" |
|
|
| from __future__ import annotations |
|
|
| from pathlib import Path |
| from typing import Any |
|
|
| import yaml |
| import qlib |
| from qlib.constant import REG_CN, REG_US |
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| DEFAULT_CONFIG = PROJECT_ROOT / "config" / "base.yaml" |
|
|
| _REGION_MAP = { |
| "cn": REG_CN, |
| "us": REG_US, |
| } |
|
|
|
|
| def load_config(config_path: str | Path | None = None) -> dict[str, Any]: |
| path = Path(config_path) if config_path else DEFAULT_CONFIG |
| with open(path, encoding="utf-8") as f: |
| cfg = yaml.safe_load(f) |
|
|
| provider_uri = cfg["qlib"]["provider_uri"] |
| provider_path = Path(provider_uri) |
| if not provider_path.is_absolute(): |
| provider_path = PROJECT_ROOT / provider_path |
| cfg["qlib"]["provider_uri_resolved"] = str(provider_path) |
| return cfg |
|
|
|
|
| def init_qlib(config_path: str | Path | None = None, exp_manager: dict | None = None) -> dict[str, Any]: |
| """Initialize qlib with project config and return the loaded config dict.""" |
| cfg = load_config(config_path) |
| region = _REGION_MAP.get(cfg["qlib"].get("region", "cn"), REG_CN) |
|
|
| freq = cfg.get("data", {}).get("freq", "day") |
| provider_uri = cfg["qlib"]["provider_uri_resolved"] |
| |
| if freq != "day": |
| provider_uri = {freq: provider_uri} |
|
|
| init_kwargs: dict[str, Any] = { |
| "provider_uri": provider_uri, |
| "region": region, |
| } |
| if exp_manager is not None: |
| init_kwargs["exp_manager"] = exp_manager |
|
|
| qlib.init(**init_kwargs) |
| return cfg |
|
|