File size: 2,768 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 | """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
|