| """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 |
|
|