"""Run qlib workflow YAML with project path injection and Recorder integration.""" from __future__ import annotations import copy import os import sys from pathlib import Path from typing import Any import yaml import qlib from qlib.model.trainer import task_train PROJECT_ROOT = Path(__file__).resolve().parents[1] def _resolve_path(value: str | Path) -> str: p = Path(value) if not p.is_absolute(): p = PROJECT_ROOT / p return str(p.resolve()) def _patch_config(config: dict[str, Any], run_id: str | None = None) -> dict[str, Any]: cfg = copy.deepcopy(config) run_id = run_id or os.environ.get("RUN_ID", "qlib_gp_run_0") qlib_init = cfg.setdefault("qlib_init", {}) if "provider_uri" in qlib_init: qlib_init["provider_uri"] = _resolve_path(qlib_init["provider_uri"]) exp_manager = qlib_init.get("exp_manager", {}) exp_kwargs = exp_manager.get("kwargs", {}) if "uri" in exp_kwargs: uri = exp_kwargs["uri"] if not str(uri).startswith("file://"): exp_kwargs["uri"] = f"file://{_resolve_path(uri)}" task = cfg.get("task", {}) dataset = task.get("dataset", {}) handler = dataset.get("kwargs", {}).get("handler", {}) handler_kwargs = handler.get("kwargs", {}) if "handler_path" in handler_kwargs or handler.get("class") == "GPFactorHandler": handler_path = PROJECT_ROOT / "outputs" / "gp_mining" / run_id / "gp_qlib_handler.pkl" if "handler_path" in handler_kwargs: raw = str(handler_kwargs["handler_path"]) if "gp_mining/" in raw: suffix = raw.split("gp_mining/", 1)[1] handler_path = PROJECT_ROOT / "outputs" / "gp_mining" / run_id / suffix.split("/", 1)[-1] else: handler_path = Path(raw) handler_kwargs["handler_path"] = _resolve_path(handler_path) return cfg def run_workflow( config_path: str | Path, experiment_name: str | None = None, run_id: str | None = None, ) -> Any: os.environ.setdefault("MLFLOW_ALLOW_FILE_STORE", "true") if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) config_path = Path(config_path) if not config_path.is_absolute(): config_path = PROJECT_ROOT / config_path with open(config_path, encoding="utf-8") as f: config = yaml.safe_load(f) config = _patch_config(config, run_id=run_id) qlib.init(**config.get("qlib_init", {})) exp_name = experiment_name or config.get("experiment_name", "workflow") recorder = task_train(config.get("task"), experiment_name=exp_name) recorder.save_objects(config=config) print(f"Experiment finished. Recorder id={recorder.id}, experiment={exp_name}") return recorder