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