1824 lines
72 KiB
Python
1824 lines
72 KiB
Python
"""TradeAC R&D explain — read mlruns experiments/runs/artifacts for the R&D UI.
|
||
|
||
This module adds the ``rd_exp_*`` MCP tool category (experiments list, per-run input config,
|
||
metrics/evaluation, model + LightGBM tree dumps, hypothesis/evaluation notes) to the
|
||
tac-qlib rd_server.
|
||
|
||
All tools are pure readers over the local MLflow **sqlite** store + artifact files:
|
||
|
||
mlruns.db -> experiments / runs / tags / params / metrics / latest_metrics
|
||
mlruns/<experiment_id>/<run_uuid>/artifacts/ -> config, params.pkl, pred.pkl, label.pkl,
|
||
sig_analysis/{ic,ric,long_short_r,long_avg_r}.pkl,
|
||
portfolio_analysis/{report,port_analysis,indicators}_*.pkl
|
||
|
||
Artifacts are Python pickles (pandas / qlib / lightgbm objects), so this module is
|
||
deliberately Python-side (see kb/guide/rd-explain.md for the architecture rationale).
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import datetime as _dt
|
||
import json
|
||
import pickle
|
||
import re
|
||
import shutil
|
||
import sqlite3
|
||
import sys
|
||
from pathlib import Path
|
||
from typing import Any, Dict, List, Optional, Tuple
|
||
|
||
import numpy as np
|
||
import pandas as pd
|
||
|
||
DEFAULT_URI = ""
|
||
DEFAULT_DEPTH = 5
|
||
DEFAULT_MAX_NODES = 256
|
||
|
||
|
||
def _default_tracking_uri() -> str:
|
||
"""Default MLflow tracking URI: Postgres ($DATABASE_URL) or lake sqlite.
|
||
|
||
With ``$DATABASE_URL`` set the store is MLflow's Postgres backend (its own
|
||
``experiments``/``runs``/... tables); artifacts live under ``<lake>/mlruns``.
|
||
Without it, the legacy unified lake sqlite store
|
||
``sqlite:///<lake>/mlruns.db`` is used. ``MLRUNS_URI`` overrides.
|
||
"""
|
||
import os
|
||
|
||
url = (os.environ.get("DATABASE_URL") or "").strip()
|
||
if url.startswith("postgres://"):
|
||
return "postgresql+psycopg://" + url[len("postgres://") :]
|
||
if url.startswith("postgresql://"):
|
||
return "postgresql+psycopg://" + url[len("postgresql://") :]
|
||
try:
|
||
from tac_qlib.data.config import resolve_lake_root
|
||
|
||
return f"sqlite:///{resolve_lake_root(None)}/mlruns.db"
|
||
except Exception: # noqa: BLE001 - fall back to the legacy repo-root store
|
||
print(
|
||
"[rd_explain] WARNING: DATABASE_URL not set — reading the lake sqlite tracking store. "
|
||
"Traced runs live in Postgres when DATABASE_URL is set.",
|
||
file=sys.stderr,
|
||
)
|
||
return "sqlite:///mlruns.db"
|
||
|
||
|
||
# --------------------------------------------------------------------------- tracking store
|
||
class _TrackingStore:
|
||
"""Dialect-agnostic MLflow tracking store (sqlite or postgres).
|
||
|
||
Exposes ``execute(sql, params)`` returning a cursor-like object with
|
||
``fetchall()``/``fetchone()`` so the rest of this module's queries are
|
||
unchanged; ``?`` placeholders are translated to ``%s`` for postgres.
|
||
"""
|
||
|
||
def __init__(self, conn, *, pg: bool, uri: str, db: str, art_root: Path):
|
||
self._conn = conn
|
||
self.pg = pg
|
||
self.uri = uri
|
||
self.db = db
|
||
self.art_root = art_root
|
||
|
||
def execute(self, sql: str, params=()):
|
||
if self.pg:
|
||
cur = self._conn.cursor()
|
||
cur.execute(sql.replace("?", "%s"), tuple(params))
|
||
return cur
|
||
return self._conn.execute(sql, params)
|
||
|
||
def close(self) -> None:
|
||
if self.pg:
|
||
self._conn.commit()
|
||
self._conn.close()
|
||
else:
|
||
self._conn.close()
|
||
|
||
|
||
def _postgres_dsn(uri: str) -> str:
|
||
"""Normalise a postgres tracking uri to a psycopg/libpq conninfo string.
|
||
|
||
MLflow URIs may be ``postgresql+psycopg://``, ``postgresql://`` or
|
||
``postgres://`` (the latter is not valid libpq); strip the SQLAlchemy
|
||
driver suffix and normalise ``postgres://`` -> ``postgresql://``.
|
||
"""
|
||
for prefix in ("postgresql+psycopg://", "postgresql://"):
|
||
if uri.startswith(prefix):
|
||
return "postgresql://" + uri[len(prefix) :]
|
||
if uri.startswith("postgres://"):
|
||
return "postgresql://" + uri[len("postgres://") :]
|
||
return uri
|
||
|
||
|
||
def _postgres_art_root() -> Path:
|
||
"""Artifact root for the postgres store: ``<lake>/mlruns`` (files stay local)."""
|
||
try:
|
||
from tac_qlib.data.config import resolve_lake_root
|
||
|
||
return resolve_lake_root(None) / "mlruns"
|
||
except Exception: # noqa: BLE001
|
||
return Path.cwd() / "mlruns"
|
||
|
||
|
||
def _sqlite_paths(uri: str) -> Tuple[Path, Path]:
|
||
"""Resolve (db_path, artifact_root) from a sqlite MLflow uri."""
|
||
if uri.startswith("sqlite:///"):
|
||
db = Path(uri[len("sqlite:///") :])
|
||
elif uri.startswith("sqlite://"):
|
||
db = Path(uri[len("sqlite://") :])
|
||
else:
|
||
db = Path(uri)
|
||
db = db.expanduser()
|
||
if not db.is_absolute():
|
||
db = Path.cwd() / db
|
||
return db, db.parent / "mlruns"
|
||
|
||
|
||
def _open_store(uri: str) -> _TrackingStore:
|
||
"""Open a tracking store for an MLflow uri (postgres or sqlite)."""
|
||
uri = (uri or _default_tracking_uri()).strip()
|
||
if uri.startswith(("postgres://", "postgresql://", "postgresql+psycopg://")):
|
||
import psycopg
|
||
|
||
conn = psycopg.connect(_postgres_dsn(uri))
|
||
return _TrackingStore(conn, pg=True, uri=uri, db=uri, art_root=_postgres_art_root())
|
||
|
||
db, art_root = _sqlite_paths(uri)
|
||
if not db.exists():
|
||
raise FileNotFoundError(f"mlruns sqlite db not found: {db}")
|
||
conn = sqlite3.connect(db)
|
||
return _TrackingStore(conn, pg=False, uri=uri, db=str(db), art_root=art_root)
|
||
|
||
|
||
def _rows(store: _TrackingStore, sql: str, params=()) -> List[tuple]:
|
||
return store.execute(sql, params).fetchall()
|
||
|
||
|
||
def _run_artifact_root(art_root: Path, exp_id, run_id: str) -> Path:
|
||
# MLflow records artifact_uri in runs; fall back to the standard layout.
|
||
return art_root / str(exp_id) / run_id / "artifacts"
|
||
|
||
|
||
def _artifact_root_for_run(store: _TrackingStore, exp_id, run_id: str) -> Path:
|
||
"""Resolve the mlflow artifact root (the ``mlruns`` dir) for a specific run.
|
||
|
||
The authoritative source is the run's recorded ``artifact_uri`` (e.g.
|
||
``<mlruns_root>/<exp>/<run>/artifacts``). Prefer the recorded location when it
|
||
matches the ``<root>/<exp>/<run>/artifacts`` layout; otherwise fall back to
|
||
the store's ``art_root``.
|
||
"""
|
||
art_root = store.art_root
|
||
meta = _run_meta(store, run_id) or {}
|
||
uri = (meta.get("artifact_uri") or "").strip()
|
||
if uri.startswith("file://"):
|
||
uri = uri[len("file://") :]
|
||
p = Path(uri).expanduser()
|
||
if p.name == "artifacts" and p.parent.name == str(run_id):
|
||
# strip <exp>/<run>/artifacts -> mlruns root; only trust it if the
|
||
# run artifact dir actually exists there (mlflow may have written it
|
||
# under a different cwd, and files can be moved after the fact).
|
||
recorded_root = p.parent.parent.parent
|
||
if (recorded_root / str(exp_id) / str(run_id) / "artifacts").exists():
|
||
return recorded_root
|
||
return art_root
|
||
|
||
|
||
# --------------------------------------------------------------------------- json utils
|
||
def _jsonable(obj: Any) -> Any:
|
||
"""Recursively convert numpy/pandas/datetime values to JSON-safe python types."""
|
||
if isinstance(obj, dict):
|
||
return {k: _jsonable(v) for k, v in obj.items()}
|
||
if isinstance(obj, (list, tuple, set)):
|
||
return [_jsonable(v) for v in obj]
|
||
if isinstance(obj, np.generic):
|
||
return obj.item()
|
||
if isinstance(obj, np.ndarray):
|
||
return obj.tolist()
|
||
if isinstance(obj, (pd.Timestamp, _dt.datetime, _dt.date, _dt.time)):
|
||
return obj.isoformat()
|
||
if isinstance(obj, float):
|
||
return obj if not np.isnan(obj) and not np.isinf(obj) else None
|
||
if obj is None or isinstance(obj, (str, int, bool)):
|
||
return obj
|
||
return str(obj)
|
||
|
||
|
||
# --------------------------------------------------------------------------- tracking store
|
||
|
||
def _list_artifact_files(root: Path, prefix: str = "") -> List[dict]:
|
||
if not root.exists():
|
||
return []
|
||
out = []
|
||
for p in sorted(root.rglob("*")):
|
||
if p.is_file():
|
||
out.append(
|
||
{
|
||
"path": str(p.relative_to(root)) if prefix else p.name,
|
||
"size": p.stat().st_size,
|
||
}
|
||
)
|
||
return out
|
||
|
||
|
||
def _load_pickle(path: Path) -> Tuple[Any, Optional[str]]:
|
||
"""Unpickle a qlib/mlflow artifact. Returns (object, error)."""
|
||
try:
|
||
with open(path, "rb") as fh:
|
||
return pickle.load(fh), None
|
||
except Exception as exc: # noqa: BLE001 - surface any unpickling failure as a warning
|
||
return None, f"unpickle {path.name}: {exc}"
|
||
|
||
|
||
def _load_series(path: Path) -> Optional[pd.Series]:
|
||
obj, err = _load_pickle(path)
|
||
if err:
|
||
return None
|
||
if isinstance(obj, pd.Series):
|
||
return obj
|
||
if isinstance(obj, pd.DataFrame) and obj.shape[1] >= 1:
|
||
return obj.iloc[:, 0]
|
||
return None
|
||
|
||
|
||
def _date_of(ts: Any) -> str:
|
||
return str(pd.Timestamp(ts).date())
|
||
|
||
|
||
# --------------------------------------------------------------------------- db helpers
|
||
def _experiment_row(conn: _TrackingStore, experiment_id) -> Optional[dict]:
|
||
rows = _rows(
|
||
conn,
|
||
"SELECT experiment_id, name, artifact_location, lifecycle_stage, creation_time, "
|
||
"last_update_time FROM experiments WHERE experiment_id = ?",
|
||
(int(experiment_id),),
|
||
)
|
||
if not rows:
|
||
return None
|
||
r = rows[0]
|
||
return {
|
||
"experiment_id": r[0],
|
||
"name": r[1],
|
||
"artifact_location": r[2],
|
||
"lifecycle_stage": r[3],
|
||
"creation_time": r[4],
|
||
"last_update_time": r[5],
|
||
}
|
||
|
||
|
||
def _run_meta(conn: _TrackingStore, run_id: str) -> Optional[dict]:
|
||
rows = _rows(
|
||
conn,
|
||
"SELECT run_uuid, name, source_type, source_name, status, start_time, end_time, "
|
||
"source_version, lifecycle_stage, artifact_uri, experiment_id FROM runs WHERE run_uuid = ?",
|
||
(run_id,),
|
||
)
|
||
if not rows:
|
||
return None
|
||
r = rows[0]
|
||
return {
|
||
"run_id": r[0],
|
||
"name": r[1],
|
||
"source_type": r[2],
|
||
"source_name": r[3],
|
||
"status": r[4],
|
||
"start_time": r[5],
|
||
"end_time": r[6],
|
||
"git_commit": r[7],
|
||
"lifecycle_stage": r[8],
|
||
"artifact_uri": r[9],
|
||
"experiment_id": r[10],
|
||
}
|
||
|
||
|
||
def _tags_of(conn: _TrackingStore, run_id: str) -> Dict[str, str]:
|
||
return {k: v for k, v in _rows(conn, "SELECT key, value FROM tags WHERE run_uuid = ?", (run_id,))}
|
||
|
||
|
||
def _params_of(conn: _TrackingStore, run_id: str) -> Dict[str, str]:
|
||
return {k: v for k, v in _rows(conn, "SELECT key, value FROM params WHERE run_uuid = ?", (run_id,))}
|
||
|
||
|
||
def _latest_metrics_of(conn: _TrackingStore, run_id: str) -> Dict[str, Any]:
|
||
rows = _rows(
|
||
conn,
|
||
"SELECT key, value FROM latest_metrics WHERE run_uuid = ? ORDER BY key",
|
||
(run_id,),
|
||
)
|
||
return {k: v for k, v in rows}
|
||
|
||
|
||
def _step_metrics_of(conn: _TrackingStore, run_id: str) -> Dict[str, list]:
|
||
rows = _rows(
|
||
conn,
|
||
"SELECT key, value, step, timestamp FROM metrics WHERE run_uuid = ? ORDER BY key, step",
|
||
(run_id,),
|
||
)
|
||
out: Dict[str, list] = {}
|
||
for k, v, step, ts in rows:
|
||
out.setdefault(k, []).append({"step": step, "value": v, "timestamp": ts})
|
||
return out
|
||
|
||
|
||
def _headline(metrics: Dict[str, Any]) -> Dict[str, Any]:
|
||
return {k: metrics[k] for k in ("IC", "ICIR", "Rank IC", "Rank ICIR") if k in metrics}
|
||
|
||
|
||
# --------------------------------------------------------------------------- notes
|
||
def _notes_path(art_root: Path, exp_id, run_id: str) -> Path:
|
||
return art_root / str(exp_id) / run_id / "rd-notes.json"
|
||
|
||
|
||
def _load_notes(art_root: Path, exp_id, run_id: str) -> dict:
|
||
p = _notes_path(art_root, exp_id, run_id)
|
||
if not p.exists():
|
||
return {"hypothesis": "", "evaluation": "", "updated_at": None}
|
||
try:
|
||
return json.loads(p.read_text())
|
||
except Exception: # noqa: BLE001
|
||
return {"hypothesis": "", "evaluation": "", "updated_at": None}
|
||
|
||
|
||
# --------------------------------------------------------------------------- input config
|
||
def _resolve_universe_from_config(hkw: Dict[str, Any], qlib_init: Dict[str, Any]) -> Tuple[List[str], str]:
|
||
"""Resolve ``instruments`` ("all"/list) to concrete symbols when the lake is available."""
|
||
instruments = hkw.get("instruments")
|
||
if isinstance(instruments, list):
|
||
return [str(s).upper() for s in instruments], "explicit"
|
||
if instruments not in (None, "all", "ALL"):
|
||
return [str(s).strip().upper() for s in str(instruments).split(",") if str(s).strip()], "explicit"
|
||
lake_root = hkw.get("lake_root") or (qlib_init or {}).get("provider_uri")
|
||
market = str(hkw.get("market") or (qlib_init or {}).get("region") or "US").upper()
|
||
try:
|
||
from tac_qlib.data.config import LakeConfig
|
||
|
||
symbols = LakeConfig(lake_root, market).load_symbols()
|
||
return symbols, "lake"
|
||
except Exception as exc: # noqa: BLE001
|
||
return [], f"unresolved ({exc})"
|
||
|
||
|
||
def _feature_fields_from_config(hkw: Dict[str, Any]) -> List[str]:
|
||
freq = hkw.get("freq", "day")
|
||
features = hkw.get("feature_fields") or hkw.get("features")
|
||
try:
|
||
from tac_qlib.contrib.data.handler import TACHandler
|
||
from tac_qlib.data.config import resolve_lake_root
|
||
|
||
return TACHandler._normalize_feature_fields(features, freq, str(resolve_lake_root(None)), "US")
|
||
except Exception: # noqa: BLE001 - feature list normalization needs the lake for defaults
|
||
if isinstance(features, list):
|
||
return list(features)
|
||
if isinstance(features, str):
|
||
return [f.strip() for f in features.split(",") if f.strip()]
|
||
return []
|
||
|
||
|
||
def _run_input(art_root: Path, exp_id, run_id: str) -> dict:
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
config = None
|
||
err = None
|
||
if (root / "config").exists():
|
||
config, err = _load_pickle(root / "config")
|
||
elif (root / "config.json").exists():
|
||
try:
|
||
config = json.loads((root / "config.json").read_text())
|
||
except Exception as exc: # noqa: BLE001
|
||
err = f"config.json: {exc}"
|
||
if config is None:
|
||
out: Dict[str, Any] = {
|
||
"source": "partial" if err else "reconstructed",
|
||
"warning": err or "no `config` artifact recorded on this run; input is not fully recoverable",
|
||
"config": None,
|
||
}
|
||
# Best-effort model settings recovered from the saved fitted model.
|
||
model, _ = _load_model(art_root, exp_id, run_id)
|
||
if model is not None:
|
||
out["model"] = {
|
||
"class": type(model).__name__,
|
||
"module_path": None,
|
||
"kwargs": _jsonable(_model_params(model)),
|
||
}
|
||
return out
|
||
|
||
task = config.get("task", {}) or {}
|
||
model = task.get("model", {}) or {}
|
||
dataset = task.get("dataset", {}) or {}
|
||
dkw = dataset.get("kwargs", {}) or {}
|
||
handler = dkw.get("handler", {}) or {}
|
||
hkw = handler.get("kwargs", {}) or {}
|
||
qlib_init = config.get("qlib_init", {}) or {}
|
||
|
||
universe, universe_source = _resolve_universe_from_config(hkw, qlib_init)
|
||
features = _feature_fields_from_config(hkw)
|
||
|
||
return {
|
||
"source": "config-artifact",
|
||
"config": _jsonable(config),
|
||
"qlib_init": _jsonable(qlib_init),
|
||
"model": _jsonable(model),
|
||
"dataset": {
|
||
"class": dataset.get("class"),
|
||
"module_path": dataset.get("module_path"),
|
||
"handler": _jsonable(handler),
|
||
"handler_kwargs": _jsonable(hkw),
|
||
"segments": _jsonable(dkw.get("segments")),
|
||
"features": features,
|
||
"label": hkw.get("label"),
|
||
"universe": universe,
|
||
"universe_size": len(universe),
|
||
"universe_source": universe_source,
|
||
},
|
||
"record": _jsonable(task.get("record")),
|
||
}
|
||
|
||
|
||
# --------------------------------------------------------------------------- result / eval
|
||
def _ic_series(art_root: Path, exp_id, run_id: str, pred=None, label=None) -> Tuple[List[dict], Optional[str]]:
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
ic = _load_series(root / "sig_analysis" / "ic.pkl")
|
||
ric = _load_series(root / "sig_analysis" / "ric.pkl")
|
||
if ic is None:
|
||
ic = _load_series(root / "ic.pkl")
|
||
if ric is None:
|
||
ric = _load_series(root / "ric.pkl")
|
||
if ic is not None and ric is not None:
|
||
out = []
|
||
for ts, i in ic.items():
|
||
out.append({"date": _date_of(ts), "ic": _jsonable(float(i)), "ric": _jsonable(float(ric.get(ts, np.nan)))})
|
||
return out, None
|
||
# Fallback: compute per-day IC/Rank IC from pred/label.
|
||
if pred is None or label is None:
|
||
return [], "ic.pkl / ric.pkl missing and pred/label unavailable to recompute"
|
||
try:
|
||
from qlib.contrib.eva.alpha import calc_ic
|
||
|
||
i2, r2 = calc_ic(pred, label, dropna=True)
|
||
out = []
|
||
for ts, i in i2.items():
|
||
out.append({"date": _date_of(ts), "ic": _jsonable(float(i)), "ric": _jsonable(float(r2.get(ts, np.nan)))})
|
||
return out, None
|
||
except Exception as exc: # noqa: BLE001
|
||
return [], f"failed to recompute IC: {exc}"
|
||
|
||
|
||
def _group_analysis(pred, label, n_groups: int = 5) -> Tuple[List[dict], Optional[pd.Series], Optional[pd.Series]]:
|
||
"""Quantile-group return analysis, mirroring qlib ``analysis_model._group_return``.
|
||
|
||
Sorts the cross-section by prediction each day, splits into ``n_groups`` blocks and
|
||
averages the label inside each block. Returns (cumulative per-group series incl.
|
||
long-short / long-average, per-day long-short, per-day long-average).
|
||
"""
|
||
if pred is None or label is None:
|
||
return [], None, None
|
||
try:
|
||
df = pd.concat([pred.rename("score"), label.rename("label")], axis=1).dropna(subset=["score", "label"])
|
||
except Exception: # noqa: BLE001 - best-effort analysis
|
||
return [], None, None
|
||
if df.empty or df.index.nlevels < 2:
|
||
return [], None, None
|
||
df = df.sort_values("score", ascending=False)
|
||
per_day: Dict[str, pd.Series] = {}
|
||
for grp in range(1, n_groups + 1):
|
||
per_day[f"Group{grp}"] = df.groupby(level=0, group_keys=False)["label"].apply(
|
||
lambda x, g=grp: x.iloc[len(x) // n_groups * (g - 1) : len(x) // n_groups * g].mean()
|
||
)
|
||
all_mean = df.groupby(level=0)["label"].mean()
|
||
per_day["long-short"] = per_day["Group1"] - per_day[f"Group{n_groups}"]
|
||
per_day["long-average"] = per_day["Group1"] - all_mean
|
||
|
||
groups: List[dict] = []
|
||
for key, s in per_day.items():
|
||
cum = s.cumsum().dropna()
|
||
groups.append(
|
||
{
|
||
"key": key,
|
||
"label": "Long-Short" if key == "long-short" else "Long-Average" if key == "long-average" else f"Group {key[5:]}",
|
||
"points": [{"date": _date_of(ts), "value": _jsonable(float(v))} for ts, v in cum.items()],
|
||
}
|
||
)
|
||
return groups, per_day["long-short"].dropna(), per_day["long-average"].dropna()
|
||
|
||
|
||
def _histogram(s, n_bins: int = 20) -> List[dict]:
|
||
"""Simple bin counts for a value series -> [{bin, count}]."""
|
||
s = pd.Series(s).dropna()
|
||
if len(s) < 2 or s.nunique() < 2:
|
||
return []
|
||
counts, edges = np.histogram(s.to_numpy(), bins=n_bins)
|
||
return [
|
||
{"bin": float((edges[i] + edges[i + 1]) / 2), "count": int(c)}
|
||
for i, c in enumerate(counts)
|
||
if c > 0
|
||
]
|
||
|
||
|
||
def _group_returns(art_root: Path, exp_id, run_id: str, pred=None, label=None) -> Tuple[dict, Optional[str]]:
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
ls = _load_series(root / "sig_analysis" / "long_short_r.pkl")
|
||
la = _load_series(root / "sig_analysis" / "long_avg_r.pkl")
|
||
if ls is None:
|
||
ls = _load_series(root / "long_short_r.pkl")
|
||
if la is None:
|
||
la = _load_series(root / "long_avg_r.pkl")
|
||
|
||
groups, ls_daily, la_daily = _group_analysis(pred, label)
|
||
distribution = {
|
||
"long_short": _histogram(ls_daily),
|
||
"long_average": _histogram(la_daily),
|
||
}
|
||
|
||
# Per-day long-short / long-average (frontend cumsums them); prefer the recorded
|
||
# artifacts and fall back to the on-the-fly computation from pred/label.
|
||
ls_pts = [{"date": _date_of(ts), "value": _jsonable(float(v))} for ts, v in ls.items()] if ls is not None else []
|
||
la_pts = [{"date": _date_of(ts), "value": _jsonable(float(v))} for ts, v in la.items()] if la is not None else []
|
||
if not ls_pts and ls_daily is not None:
|
||
ls_pts = [{"date": _date_of(ts), "value": _jsonable(float(v))} for ts, v in ls_daily.items()]
|
||
if not la_pts and la_daily is not None:
|
||
la_pts = [{"date": _date_of(ts), "value": _jsonable(float(v))} for ts, v in la_daily.items()]
|
||
|
||
warn = None
|
||
if ls is None and la is None and (pred is None or label is None):
|
||
warn = "long_short_r.pkl / long_avg_r.pkl missing and pred/label unavailable to recompute"
|
||
return {
|
||
"long_short": ls_pts,
|
||
"long_avg": la_pts,
|
||
"groups": groups,
|
||
"distribution": distribution,
|
||
}, warn
|
||
|
||
|
||
def _monthly_ic(ic_series: List[dict]) -> List[dict]:
|
||
"""Aggregate the daily IC / Rank IC series into a year x month grid."""
|
||
rows: Dict[str, Dict[str, list]] = {}
|
||
for p in ic_series:
|
||
ym = p["date"][:7]
|
||
row = rows.setdefault(ym, {"ic": [], "ric": []})
|
||
if p.get("ic") is not None:
|
||
row["ic"].append(p["ic"])
|
||
if p.get("ric") is not None:
|
||
row["ric"].append(p["ric"])
|
||
out = []
|
||
for ym in sorted(rows):
|
||
r = rows[ym]
|
||
out.append(
|
||
{
|
||
"month": ym,
|
||
"year": int(ym[:4]),
|
||
"month_of_year": int(ym[5:7]),
|
||
"ic": round(float(np.mean(r["ic"])), 4) if r["ic"] else None,
|
||
"ric": round(float(np.mean(r["ric"])), 4) if r["ric"] else None,
|
||
}
|
||
)
|
||
return out
|
||
|
||
|
||
def _ic_histogram(ic_series: List[dict]) -> dict:
|
||
ic = [p["ic"] for p in ic_series if p.get("ic") is not None]
|
||
ric = [p["ric"] for p in ic_series if p.get("ric") is not None]
|
||
return {"ic": _histogram(pd.Series(ic)), "ric": _histogram(pd.Series(ric))}
|
||
|
||
|
||
def _autocorrelation(pred) -> List[dict]:
|
||
"""Per-day rank autocorrelation of the prediction (lag 1), mirroring qlib ``_pred_autocorr``."""
|
||
if pred is None:
|
||
return []
|
||
s = pred.dropna()
|
||
if s.empty or s.index.nlevels < 2:
|
||
return []
|
||
df = s.to_frame("score")
|
||
df["score_last"] = df.groupby(level=1, group_keys=False)["score"].shift(1)
|
||
ac = df.groupby(level=0, group_keys=False).apply(
|
||
lambda x: x["score"].rank(pct=True).corr(x["score_last"].rank(pct=True))
|
||
)
|
||
out = []
|
||
for ts, v in ac.items():
|
||
if v is not None and np.isfinite(v):
|
||
out.append({"date": _date_of(ts), "value": _jsonable(float(v))})
|
||
return out
|
||
|
||
|
||
def _explanation(headline: Dict[str, Any], risk: Dict[str, Any], ic_series: List[dict]) -> List[str]:
|
||
"""Auto-generated findings interpreting the headline / risk metrics and daily IC."""
|
||
lines: List[str] = []
|
||
ic = headline.get("IC")
|
||
icir = headline.get("ICIR")
|
||
if ic is not None:
|
||
icir_txt = f" (ICIR {icir:+.2f})" if icir is not None else ""
|
||
lines.append(f"Average daily IC {ic:+.4f}{icir_txt} across the test window.")
|
||
if icir is not None:
|
||
if abs(icir) >= 0.5:
|
||
lines.append("ICIR above 0.5 in magnitude — the signal's predictive power is unlikely to be pure luck.")
|
||
elif abs(icir) >= 0.2:
|
||
lines.append("ICIR in the 0.2–0.5 range — weak but persistent predictive power.")
|
||
else:
|
||
lines.append("ICIR below 0.2 in magnitude — the signal is barely distinguishable from noise.")
|
||
if ic_series:
|
||
valid = [p["ic"] for p in ic_series if p.get("ic") is not None]
|
||
if valid:
|
||
pos = sum(1 for v in valid if v > 0) / len(valid)
|
||
lines.append(f"IC was positive on {pos * 100:.0f}% of days ({len(valid)} days).")
|
||
ann = risk.get("annualized_return") or risk.get("1day.annualized_return")
|
||
mdd = risk.get("max_drawdown") or risk.get("1day.max_drawdown")
|
||
sharpe = risk.get("information_ratio") or risk.get("1day.information_ratio")
|
||
if ann is not None:
|
||
lines.append(f"Backtest annualized return {ann * 100:+.1f}% with max drawdown {mdd * 100:.1f}%." if mdd is not None else f"Backtest annualized return {ann * 100:+.1f}%.")
|
||
if sharpe is not None:
|
||
lines.append(f"Information ratio (Sharpe) {sharpe:.2f}.")
|
||
return lines
|
||
|
||
|
||
def _flatten_risk(df: pd.DataFrame) -> Dict[str, Any]:
|
||
"""Normalize a risk-analysis DataFrame to flat metric keys.
|
||
|
||
qlib's ``port_analysis_*.pkl`` rows may carry a MultiIndex like
|
||
``(excess_return_with_cost, annualized_return)``. We emit both prefixed keys
|
||
(``excess_return_with_cost.annualized_return``) and, for the preferred
|
||
``excess_return_with_cost`` analysis, bare metric names so the UI and the
|
||
explanation generator can read ``annualized_return`` / ``max_drawdown`` /
|
||
``information_ratio`` directly.
|
||
"""
|
||
col = df.columns[0]
|
||
out: Dict[str, Any] = {}
|
||
for idx, row in df.iterrows():
|
||
val = row[col]
|
||
try:
|
||
val = _jsonable(float(val))
|
||
except (TypeError, ValueError): # noqa: S110 - non-numeric rows are skipped
|
||
continue
|
||
if isinstance(idx, tuple):
|
||
analysis, metric = idx
|
||
out[f"{analysis}.{metric}"] = val
|
||
if analysis == "excess_return_with_cost" and metric not in out:
|
||
out[metric] = val
|
||
else:
|
||
out[str(idx)] = val
|
||
return out
|
||
|
||
|
||
def _backtest(art_root: Path, exp_id, run_id: str) -> Tuple[List[dict], Dict[str, Any], Optional[str], Optional[str], Dict[str, Any]]:
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
warn = None
|
||
report = None
|
||
for cand in (
|
||
"portfolio_analysis/report_normal_1day.pkl",
|
||
"portfolio_analysis/report_normal_1d.pkl",
|
||
"report_normal_1day.pkl",
|
||
"report_normal_1d.pkl",
|
||
):
|
||
p = root / cand
|
||
if p.exists():
|
||
obj, err = _load_pickle(p)
|
||
if err:
|
||
warn = err
|
||
elif isinstance(obj, pd.DataFrame) and len(obj):
|
||
report = obj
|
||
break
|
||
|
||
pts: List[dict] = []
|
||
has_bench = False
|
||
ann: Dict[str, Any] = {}
|
||
if report is not None:
|
||
cum = report["return"].fillna(0).cumsum()
|
||
bench = report["bench"].fillna(0).cumsum() if "bench" in report.columns else None
|
||
has_bench = bench is not None and bool(report["bench"].abs().sum())
|
||
# qlib annualizes day-frequency returns with the 238 scaler (see
|
||
# qlib.contrib.evaluate.risk_analysis, sum mode): mean * 238.
|
||
ann["return_annualized"] = _jsonable(float(report["return"].fillna(0).mean() * 238))
|
||
ann["benchmark_annualized"] = _jsonable(float(report["bench"].fillna(0).mean() * 238)) if has_bench else None
|
||
for idx, r in report.iterrows():
|
||
rec = {"date": _date_of(idx), "return": _jsonable(float(r.get("return", 0)))}
|
||
rec["cum_return"] = _jsonable(float(cum.loc[idx]))
|
||
if bench is not None:
|
||
rec["cum_benchmark"] = _jsonable(float(bench.loc[idx]))
|
||
pts.append(rec)
|
||
|
||
risk: Dict[str, Any] = {}
|
||
for cand in (
|
||
"portfolio_analysis/port_analysis_1day.pkl",
|
||
"portfolio_analysis/port_analysis_1d.pkl",
|
||
"port_analysis_1day.pkl",
|
||
"port_analysis_1d.pkl",
|
||
):
|
||
p = root / cand
|
||
if p.exists():
|
||
obj, _ = _load_pickle(p)
|
||
if isinstance(obj, pd.DataFrame) and len(obj):
|
||
risk = _flatten_risk(obj)
|
||
break
|
||
if not risk and report is not None:
|
||
try:
|
||
from qlib.contrib.evaluate import risk_analysis
|
||
|
||
ra = risk_analysis(report["return"], freq="day")
|
||
risk = {str(k): _jsonable(float(v)) for k, v in ra.iloc[:, 0].items()}
|
||
except Exception as exc: # noqa: BLE001
|
||
warn = (warn + "; " if warn else "") + f"risk_analysis failed: {exc}"
|
||
|
||
benchmark = _backtest_benchmark(art_root, exp_id, run_id)
|
||
if not benchmark and not has_bench:
|
||
benchmark = None
|
||
return pts, risk, warn, benchmark, ann
|
||
|
||
|
||
def _pred_stats(art_root: Path, exp_id, run_id: str) -> Tuple[Optional[dict], Optional[str]]:
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
pred = _load_series(root / "pred.pkl")
|
||
if pred is None:
|
||
return None, "pred.pkl missing"
|
||
s = pred.dropna()
|
||
insts = sorted(s.index.get_level_values(1).dropna().astype(str).unique()) if s.index.nlevels > 1 else []
|
||
return {
|
||
"count": int(len(s)),
|
||
"date_min": _date_of(s.index.get_level_values(0).min()),
|
||
"date_max": _date_of(s.index.get_level_values(0).max()),
|
||
"instruments": insts,
|
||
"instrument_count": len(insts),
|
||
"mean": round(float(s.mean()), 6),
|
||
"std": round(float(s.std(ddof=1)), 6) if len(s) > 1 else None,
|
||
"min": round(float(s.min()), 6),
|
||
"max": round(float(s.max()), 6),
|
||
}, None
|
||
|
||
|
||
def _strategy_kwargs(art_root: Path, exp_id, run_id: str) -> dict:
|
||
"""Recover the PortAnaRecord strategy kwargs (topk, n_drop, benchmark) from the saved config."""
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
config, _ = _load_pickle(root / "config") if (root / "config").exists() else (None, None)
|
||
if not config:
|
||
return {}
|
||
try:
|
||
for rec in config.get("task", {}).get("record", []) or []:
|
||
if str(rec.get("class", "")).endswith("PortAnaRecord"):
|
||
strategy = ((rec.get("kwargs", {}) or {}).get("config", {}) or {}).get("strategy", {}) or {}
|
||
return dict(strategy.get("kwargs", {}) or {})
|
||
except Exception: # noqa: BLE001 - best-effort recovery
|
||
pass
|
||
return {}
|
||
|
||
|
||
def _blotter_topk(art_root: Path, exp_id, run_id: str, default: int = 2) -> int:
|
||
"""Recover the strategy ``topk`` from the recorded config (PortAnaRecord)."""
|
||
topk = _strategy_kwargs(art_root, exp_id, run_id).get("topk")
|
||
return int(topk) if topk else default
|
||
|
||
|
||
def _backtest_benchmark(art_root: Path, exp_id, run_id: str) -> Optional[str]:
|
||
"""Best-effort benchmark symbol from the recorded config (may be empty for legacy runs)."""
|
||
bm = (_strategy_kwargs(art_root, exp_id, run_id).get("benchmark") or "").strip()
|
||
return bm or None
|
||
|
||
|
||
def _order_trades(root: Path) -> Tuple[List[dict], Optional[str]]:
|
||
"""Extract per-day per-symbol executed trades from the OrderIndicator artifact.
|
||
|
||
``indicators_*_obj.pkl`` holds the qlib ``OrderIndicator`` whose history maps
|
||
each trading day to per-symbol series: amount (signed shares), deal price,
|
||
trade value, trade cost and trade direction (1 buy / 0 sell).
|
||
"""
|
||
cands = (
|
||
"portfolio_analysis/indicators_normal_1day_obj.pkl",
|
||
"portfolio_analysis/indicators_1day_obj.pkl",
|
||
"portfolio_analysis/indicators_normal_1day.pkl",
|
||
"indicators_normal_1day_obj.pkl",
|
||
"indicators_1day_obj.pkl",
|
||
)
|
||
io = None
|
||
for cand in cands:
|
||
p = root / cand
|
||
if p.exists():
|
||
obj, err = _load_pickle(p)
|
||
if err:
|
||
return [], err
|
||
if obj is not None and hasattr(obj, "order_indicator_his"):
|
||
io = obj
|
||
break
|
||
if io is None:
|
||
return [], "order/trade indicator artifact missing (no backtest trades recorded)"
|
||
|
||
def _to_series(sd):
|
||
try:
|
||
return pd.Series(sd.data, index=list(sd.index))
|
||
except Exception: # noqa: BLE001 - qlib SingleData -> pandas
|
||
return pd.Series(dtype=float)
|
||
|
||
trades: List[dict] = []
|
||
for ts, v in io.order_indicator_his.items():
|
||
d = getattr(v, "data", v)
|
||
if not isinstance(d, dict):
|
||
continue
|
||
amt = _to_series(d.get("amount"))
|
||
if not len(amt):
|
||
continue
|
||
price = _to_series(d.get("trade_price"))
|
||
value = _to_series(d.get("trade_value"))
|
||
cost = _to_series(d.get("trade_cost"))
|
||
tdir = _to_series(d.get("trade_dir"))
|
||
for sym in amt.index:
|
||
a = float(amt[sym])
|
||
if abs(a) < 1e-9:
|
||
continue
|
||
trades.append(
|
||
{
|
||
"date": _date_of(ts),
|
||
"symbol": str(sym),
|
||
"side": "BUY" if float(tdir.get(sym, 1)) >= 0.5 else "SELL",
|
||
"amount": a,
|
||
"price": _jsonable(float(price.get(sym, np.nan))),
|
||
"value": _jsonable(float(value.get(sym, np.nan))),
|
||
"cost": _jsonable(float(cost.get(sym, np.nan))),
|
||
}
|
||
)
|
||
trades.sort(key=lambda t: (t["date"], t["symbol"]))
|
||
return trades, None
|
||
|
||
|
||
def _fifo_trade_pnl(trades: List[dict]) -> Tuple[float, float]:
|
||
"""FIFO tax-lot accounting for trade-level P&L.
|
||
|
||
Every BUY opens a lot (with its leading signal). Every SELL consumes the
|
||
oldest open lots first (FIFO): realized P&L = sell proceeds minus the
|
||
matched lots' all-in cost basis, and the trade inherits the entry date,
|
||
entry price and the leading signal of the matched lot(s). Returns
|
||
(realized_pnl, unrealized_pnl).
|
||
"""
|
||
lots: Dict[str, List[dict]] = {}
|
||
realized_total = 0.0
|
||
for t in trades:
|
||
sym = t["symbol"]
|
||
t.setdefault("signal_score", None)
|
||
t.setdefault("signal_rank", None)
|
||
t["entry_date"] = None
|
||
t["entry_price"] = None
|
||
t["entry_score"] = None
|
||
t["realized_pnl"] = None
|
||
if t["side"] == "BUY":
|
||
amount = float(t["amount"])
|
||
if amount <= 0:
|
||
continue
|
||
unit_cost = (float(t["value"] or 0) + float(t["cost"] or 0)) / amount
|
||
lots.setdefault(sym, []).append(
|
||
{"qty": amount, "unit_cost": unit_cost, "date": t["date"], "score": t["signal_score"]}
|
||
)
|
||
else:
|
||
qty_sell = -float(t["amount"])
|
||
proceeds = -float(t["value"] or 0) - float(t["cost"] or 0)
|
||
open_lots = lots.setdefault(sym, [])
|
||
remaining = qty_sell
|
||
basis = 0.0
|
||
matched = []
|
||
while remaining > 1e-9 and open_lots:
|
||
lot = open_lots[0]
|
||
take = min(remaining, lot["qty"])
|
||
basis += take * lot["unit_cost"]
|
||
matched.append((take, lot))
|
||
lot["qty"] -= take
|
||
remaining -= take
|
||
if lot["qty"] < 1e-9:
|
||
open_lots.pop(0)
|
||
realized = proceeds - basis
|
||
realized_total += realized
|
||
t["realized_pnl"] = round(realized, 6)
|
||
t["entry_date"] = matched[0][1]["date"] if matched else None
|
||
t["entry_score"] = matched[0][1]["score"] if matched else None
|
||
if matched and qty_sell > 0:
|
||
t["entry_price"] = round(sum(take * lot["unit_cost"] for take, lot in matched) / qty_sell, 6)
|
||
|
||
unrealized_total = 0.0
|
||
return realized_total, unrealized_total
|
||
|
||
|
||
def _blotter(art_root: Path, exp_id, run_id: str) -> Tuple[dict, List[str]]:
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
warnings: List[str] = []
|
||
|
||
report = None
|
||
for cand in (
|
||
"portfolio_analysis/report_normal_1day.pkl",
|
||
"portfolio_analysis/report_1day.pkl",
|
||
"report_normal_1day.pkl",
|
||
"report_1day.pkl",
|
||
):
|
||
p = root / cand
|
||
if p.exists():
|
||
obj, err = _load_pickle(p)
|
||
if err:
|
||
warnings.append(err)
|
||
elif isinstance(obj, pd.DataFrame) and len(obj):
|
||
report = obj
|
||
break
|
||
|
||
positions: Optional[dict] = None
|
||
for cand in (
|
||
"portfolio_analysis/positions_normal_1day.pkl",
|
||
"portfolio_analysis/positions_1day.pkl",
|
||
"positions_normal_1day.pkl",
|
||
"positions_1day.pkl",
|
||
):
|
||
p = root / cand
|
||
if p.exists():
|
||
obj, err = _load_pickle(p)
|
||
if err:
|
||
warnings.append(err)
|
||
elif isinstance(obj, dict) and len(obj):
|
||
positions = obj
|
||
break
|
||
|
||
trades, trade_warn = _order_trades(root)
|
||
if trade_warn:
|
||
warnings.append(trade_warn)
|
||
|
||
pred = _load_series(root / "pred.pkl")
|
||
topk = _blotter_topk(art_root, exp_id, run_id)
|
||
|
||
signal_map: Dict[Tuple[str, str], float] = {}
|
||
signal_rank_map: Dict[Tuple[str, str], int] = {}
|
||
signal_rows: List[dict] = []
|
||
if pred is not None and pred.index.nlevels >= 2:
|
||
df = pred.to_frame("score").dropna()
|
||
df = df.reset_index()
|
||
df.columns = ["date", "symbol", "score"][: df.shape[1]]
|
||
df["date"] = df["date"].map(_date_of)
|
||
df["rank"] = df.groupby("date")["score"].rank(ascending=False, method="first").astype(int)
|
||
traded_dates = {(t["date"], t["symbol"]) for t in trades}
|
||
for _, r in df.iterrows():
|
||
key = (str(r["date"]), str(r["symbol"]))
|
||
score = float(r["score"])
|
||
rank = int(r["rank"])
|
||
signal_map[key] = score
|
||
signal_rank_map[key] = rank
|
||
signal_rows.append(
|
||
{
|
||
"date": str(r["date"]),
|
||
"symbol": str(r["symbol"]),
|
||
"score": round(score, 6),
|
||
"rank": rank,
|
||
"selected": rank <= topk,
|
||
"traded": key in traded_dates,
|
||
}
|
||
)
|
||
elif pred is None:
|
||
warnings.append("pred.pkl missing; signal blotter unavailable")
|
||
|
||
for t in trades:
|
||
key = (t["date"], t["symbol"])
|
||
t["signal_score"] = signal_map.get(key)
|
||
t["signal_rank"] = signal_rank_map.get(key)
|
||
_fifo_trade_pnl(trades)
|
||
|
||
# Daily report series + P&L
|
||
daily: List[dict] = []
|
||
initial_account = None
|
||
final_account = None
|
||
if report is not None:
|
||
prev = None
|
||
for idx, row in report.iterrows():
|
||
acc = float(row.get("account", np.nan))
|
||
if initial_account is None and not np.isnan(acc):
|
||
initial_account = acc
|
||
pnl_day = acc - prev if prev is not None else 0.0
|
||
prev = acc
|
||
daily.append(
|
||
{
|
||
"date": _date_of(idx),
|
||
"account": round(acc, 6) if not np.isnan(acc) else None,
|
||
"value": _jsonable(float(row.get("value", np.nan))),
|
||
"cash": _jsonable(float(row.get("cash", np.nan))),
|
||
"return": _jsonable(float(row.get("return", np.nan))),
|
||
"cost_ratio": _jsonable(float(row.get("cost", np.nan))),
|
||
"pnl_day": round(pnl_day, 6),
|
||
}
|
||
)
|
||
if daily:
|
||
final_account = daily[-1]["account"]
|
||
|
||
# Final position snapshot + position history
|
||
position_history: List[dict] = []
|
||
current: List[dict] = []
|
||
final_prices: Dict[str, float] = {}
|
||
if positions is not None:
|
||
for ts, pos in positions.items():
|
||
pos_d = getattr(pos, "position", {}) or {}
|
||
for sym, info in pos_d.items():
|
||
if sym in ("cash", "now_account_value"):
|
||
continue
|
||
if not isinstance(info, dict):
|
||
continue
|
||
amount = float(info.get("amount", 0))
|
||
if abs(amount) < 1e-9:
|
||
continue
|
||
price = _jsonable(float(info.get("price", np.nan)))
|
||
position_history.append(
|
||
{"date": _date_of(ts), "symbol": str(sym), "amount": round(amount, 6), "price": price}
|
||
)
|
||
final_prices[str(sym)] = float(info.get("price", np.nan)) if info.get("price") is not None else np.nan
|
||
position_history.sort(key=lambda r: (r["date"], r["symbol"]))
|
||
|
||
# FIFO open-lot ledger -> per-symbol avg cost, unrealized P&L, entry signal.
|
||
unrealized_by_symbol: Dict[str, float] = {}
|
||
avg_cost_by_symbol: Dict[str, float] = {}
|
||
position_entry_by_symbol: Dict[str, Optional[str]] = {}
|
||
first_score_by_symbol: Dict[str, Optional[float]] = {}
|
||
realized_by_symbol: Dict[str, float] = {}
|
||
if trades:
|
||
lots: Dict[str, List[dict]] = {}
|
||
for t in trades:
|
||
sym = t["symbol"]
|
||
if t["side"] == "BUY" and float(t["amount"]) > 0:
|
||
amount = float(t["amount"])
|
||
unit_cost = (float(t["value"] or 0) + float(t["cost"] or 0)) / amount
|
||
lots.setdefault(sym, []).append(
|
||
{"qty": amount, "unit_cost": unit_cost, "date": t["date"], "score": t["signal_score"]}
|
||
)
|
||
elif t["side"] == "SELL":
|
||
open_lots = lots.setdefault(sym, [])
|
||
remaining = -float(t["amount"])
|
||
while remaining > 1e-9 and open_lots:
|
||
lot = open_lots[0]
|
||
take = min(remaining, lot["qty"])
|
||
lot["qty"] -= take
|
||
remaining -= take
|
||
if lot["qty"] < 1e-9:
|
||
open_lots.pop(0)
|
||
for sym, open_lots in lots.items():
|
||
price = final_prices.get(sym, np.nan)
|
||
qty = sum(lot["qty"] for lot in open_lots)
|
||
if qty > 1e-9:
|
||
basis = sum(lot["qty"] * lot["unit_cost"] for lot in open_lots)
|
||
avg_cost_by_symbol[sym] = basis / qty
|
||
position_entry_by_symbol[sym] = open_lots[0]["date"]
|
||
first_score_by_symbol[sym] = open_lots[0]["score"]
|
||
if np.isfinite(price):
|
||
unrealized_by_symbol[sym] = sum(
|
||
lot["qty"] * (price - lot["unit_cost"]) for lot in open_lots
|
||
)
|
||
|
||
realized_by_symbol: Dict[str, float] = {}
|
||
for t in trades:
|
||
if t["side"] == "SELL" and t["realized_pnl"] is not None:
|
||
realized_by_symbol[t["symbol"]] = realized_by_symbol.get(t["symbol"], 0.0) + t["realized_pnl"]
|
||
|
||
if positions:
|
||
last_pos = getattr(list(positions.values())[-1], "position", {}) or {}
|
||
for sym, info in last_pos.items():
|
||
if sym in ("cash", "now_account_value") or not isinstance(info, dict):
|
||
continue
|
||
amount = float(info.get("amount", 0))
|
||
if abs(amount) < 1e-9:
|
||
continue
|
||
price = float(info.get("price", np.nan)) if info.get("price") is not None else np.nan
|
||
weight = float(info.get("weight", 0)) if info.get("weight") is not None else np.nan
|
||
days_held = int(info.get("count_day", 0)) if info.get("count_day") is not None else 0
|
||
current.append(
|
||
{
|
||
"symbol": str(sym),
|
||
"amount": round(amount, 6),
|
||
"price": round(price, 6) if np.isfinite(price) else None,
|
||
"value": round(amount * price, 2) if np.isfinite(price) else None,
|
||
"weight": round(weight, 4) if np.isfinite(weight) else None,
|
||
"days_held": days_held,
|
||
"avg_cost": round(avg_cost_by_symbol.get(sym, np.nan), 6)
|
||
if sym in avg_cost_by_symbol
|
||
else None,
|
||
"position_entry_date": position_entry_by_symbol.get(sym),
|
||
"entry_score": round(first_score_by_symbol[sym], 6)
|
||
if first_score_by_symbol.get(sym) is not None
|
||
else None,
|
||
"realized_pnl": round(realized_by_symbol.get(sym, 0.0), 2),
|
||
"unrealized_pnl": round(unrealized_by_symbol.get(sym, 0.0), 2),
|
||
}
|
||
)
|
||
|
||
realized_total = round(sum(realized_by_symbol.values()), 2)
|
||
unrealized_total = round(sum(unrealized_by_symbol.values()), 2)
|
||
|
||
summary: Dict[str, Any] = {
|
||
"initial_account": round(initial_account, 2) if initial_account is not None else None,
|
||
"final_account": round(final_account, 2) if final_account is not None else None,
|
||
"pnl_inception": round(final_account - initial_account, 2)
|
||
if initial_account is not None and final_account is not None
|
||
else None,
|
||
"realized_pnl": realized_total,
|
||
"unrealized_pnl": unrealized_total,
|
||
"total_cost": round(sum(float(t.get("cost") or 0.0) for t in trades), 2),
|
||
"trading_days": len(daily),
|
||
"topk": topk,
|
||
"n_trades": len(trades),
|
||
"n_buys": sum(1 for t in trades if t["side"] == "BUY"),
|
||
"n_sells": sum(1 for t in trades if t["side"] == "SELL"),
|
||
}
|
||
|
||
return {
|
||
"summary": summary,
|
||
"daily": daily,
|
||
"positions": current,
|
||
"position_history": position_history,
|
||
"trades": trades,
|
||
"signals": signal_rows,
|
||
"warnings": warnings,
|
||
}, warnings
|
||
|
||
|
||
# --------------------------------------------------------------------------- model / tree
|
||
def _load_model(art_root: Path, exp_id, run_id: str) -> Tuple[Any, Optional[str]]:
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
p = root / "params.pkl"
|
||
if not p.exists():
|
||
return None, "params.pkl missing"
|
||
return _load_pickle(p)
|
||
|
||
|
||
def _model_params(model: Any) -> Dict[str, Any]:
|
||
"""Best-effort training params recovered from a saved fitted model."""
|
||
try:
|
||
params = dict(model.params)
|
||
except Exception: # noqa: BLE001 - params are best-effort
|
||
params = {}
|
||
if not params:
|
||
booster = getattr(model, "model", None)
|
||
if booster is None:
|
||
# Ensemble: sub-models carry `.params`; recover from the first one.
|
||
subs = getattr(model, "_models", None) or []
|
||
if subs:
|
||
try:
|
||
params = dict(subs[0].params)
|
||
except Exception: # noqa: BLE001
|
||
params = {}
|
||
if not params and booster is not None:
|
||
try:
|
||
params = dict(booster.params)
|
||
except Exception: # noqa: BLE001 - params are best-effort
|
||
params = {}
|
||
return params
|
||
|
||
|
||
_COLUMN_NAME = re.compile(r"^Column_\d+$")
|
||
|
||
|
||
def _resolve_feature_names(art_root: Path, exp_id, run_id: str, feature_names: List[str]) -> List[str]:
|
||
"""Map LightGBM's ``Column_N`` placeholders to the real handler feature fields.
|
||
|
||
qlib trains on numpy arrays, so the booster reports ``Column_0..N``. The saved config
|
||
artifact carries the resolved ``feature_fields``; when their order/length match the
|
||
booster columns, substitute the real names so the tree and importances are readable.
|
||
"""
|
||
if not feature_names or not all(_COLUMN_NAME.match(n) for n in feature_names):
|
||
return feature_names
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
config, _ = _load_pickle(root / "config") if (root / "config").exists() else (None, None)
|
||
if not config:
|
||
# Legacy runs trained before the config artifact was persisted: reconstruct the
|
||
# handler's default feature list from the current lake (raw OHLCV + common ta-lib).
|
||
try:
|
||
from tac_qlib.contrib.data.handler import TACHandler
|
||
from tac_qlib.data.config import resolve_lake_root
|
||
|
||
real = TACHandler._normalize_feature_fields(None, "day", str(resolve_lake_root(None)), "US")
|
||
if len(real) == len(feature_names):
|
||
return list(real)
|
||
except Exception: # noqa: BLE001, S110 - best-effort reconstruction
|
||
pass
|
||
return feature_names
|
||
dkw = ((config.get("task", {}) or {}).get("dataset", {}) or {}).get("kwargs", {}) or {}
|
||
hkw = (dkw.get("handler", {}) or {}).get("kwargs", {}) or {}
|
||
real = _feature_fields_from_config(hkw)
|
||
if len(real) == len(feature_names):
|
||
return list(real)
|
||
return feature_names
|
||
|
||
|
||
def _flatten_tree(node: dict, depth: int, max_depth: int, feature_names: List[str], nodes: List[dict]) -> bool:
|
||
"""Flatten one LightGBM tree node recursively (pre-order).
|
||
|
||
Each node gets ``id = len(nodes)`` at append time; children are appended right after
|
||
their parent, so ``left`` is ``id+1`` and ``right`` starts where the left subtree ends.
|
||
Returns False when the depth or node budget is hit (parent then marks that child None).
|
||
"""
|
||
if node is None or depth > max_depth or len(nodes) >= DEFAULT_MAX_NODES:
|
||
return False
|
||
idx = len(nodes)
|
||
if "leaf_value" in node:
|
||
nodes.append(
|
||
{
|
||
"id": idx,
|
||
"depth": depth,
|
||
"leaf": True,
|
||
"value": _jsonable(float(node.get("leaf_value", np.nan))),
|
||
"count": int(node.get("leaf_count", 0)),
|
||
}
|
||
)
|
||
return True
|
||
sf = node.get("split_feature")
|
||
feat = feature_names[sf] if isinstance(sf, int) and sf < len(feature_names) else str(sf)
|
||
nodes.append(
|
||
{
|
||
"id": idx,
|
||
"depth": depth,
|
||
"leaf": False,
|
||
"feature": feat,
|
||
"threshold": _jsonable(float(node.get("threshold", np.nan))),
|
||
"gain": _jsonable(float(node.get("split_gain", np.nan))),
|
||
"count": int(node.get("internal_count", 0)),
|
||
}
|
||
)
|
||
left_ok = _flatten_tree(node.get("left_child"), depth + 1, max_depth, feature_names, nodes)
|
||
nodes[idx]["left"] = idx + 1 if left_ok else None
|
||
right_start = len(nodes)
|
||
right_ok = _flatten_tree(node.get("right_child"), depth + 1, max_depth, feature_names, nodes)
|
||
nodes[idx]["right"] = right_start if right_ok else None
|
||
return True
|
||
|
||
|
||
def _model_dump(art_root: Path, exp_id, run_id: str, tree_id: int, max_depth: int, seed: int = 0) -> Tuple[Optional[dict], Optional[str]]:
|
||
model, err = _load_model(art_root, exp_id, run_id)
|
||
if err:
|
||
return None, err
|
||
if model is None:
|
||
return None, "model object missing"
|
||
booster = getattr(model, "model", None)
|
||
# Ensemble models (RankICEnsembleLGBModel) hold one fitted sub-model per
|
||
# seed in `_models` with no top-level `.model` booster. Dump a specific
|
||
# member (default seed 0) so the tree/importance view works per seed.
|
||
sub_index = None
|
||
if booster is None:
|
||
subs = getattr(model, "_models", None) or []
|
||
if subs:
|
||
sub_index = min(int(seed), len(subs) - 1)
|
||
booster = getattr(subs[sub_index], "model", None)
|
||
if booster is None or not hasattr(booster, "dump_model"):
|
||
return None, f"model type {type(model).__name__} has no LightGBM booster to dump"
|
||
|
||
try:
|
||
dump = booster.dump_model()
|
||
except Exception as exc: # noqa: BLE001
|
||
return None, f"dump_model failed: {exc}"
|
||
|
||
feature_names = _resolve_feature_names(art_root, exp_id, run_id, list(dump.get("feature_names") or []))
|
||
tree_info = dump.get("tree_info") or []
|
||
if not tree_info:
|
||
return None, "empty tree_info"
|
||
num_trees = len(tree_info)
|
||
tree = tree_info[tree_id % num_trees]
|
||
|
||
# A single tree is only weakly representative of a boosted model; expose
|
||
# "milestone" tree ids (first / middle / best-iteration) so the UI can offer
|
||
# quick jumps instead of defaulting to the least-informative tree 0.
|
||
best_it = None
|
||
for attr in ("best_iteration", "best_iter"):
|
||
v = getattr(model, attr, None)
|
||
if v is not None:
|
||
best_it = int(v)
|
||
break
|
||
if best_it is None and sub_index is not None:
|
||
subs = getattr(model, "_models", None) or []
|
||
if subs:
|
||
for attr in ("best_iteration", "best_iter"):
|
||
v = getattr(subs[sub_index], attr, None)
|
||
if v is not None:
|
||
best_it = int(v)
|
||
break
|
||
best_tree_id = (best_it - 1) if best_it else (num_trees // 2)
|
||
milestones = {
|
||
"first": 0,
|
||
"mid": max(0, num_trees // 2),
|
||
"best": max(0, min(best_tree_id, num_trees - 1)),
|
||
}
|
||
|
||
importances = {}
|
||
try:
|
||
imp = booster.feature_importance(importance_type="gain")
|
||
for i, name in enumerate(feature_names):
|
||
importances[name] = round(float(imp[i]), 4) if i < len(imp) else 0.0
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
|
||
nodes: List[dict] = []
|
||
_flatten_tree(tree.get("tree_structure"), 0, int(max_depth), feature_names, nodes)
|
||
|
||
return {
|
||
"class": type(model).__name__,
|
||
"sub_model": sub_index,
|
||
"seed_count": len(getattr(model, "_models", None) or []) or 1,
|
||
"num_trees": num_trees,
|
||
"best_iteration": best_it,
|
||
"milestone_tree_ids": milestones,
|
||
"feature_names": feature_names,
|
||
"feature_importances": importances,
|
||
"tree_id": tree_id % num_trees,
|
||
"tree_index": tree.get("tree_index", tree_id % num_trees),
|
||
"num_leaves": tree.get("num_leaves"),
|
||
"truncated": len(nodes) >= DEFAULT_MAX_NODES,
|
||
"nodes": nodes,
|
||
}, None
|
||
|
||
|
||
# --------------------------------------------------------------------------- tools
|
||
def rd_exp_list(uri: str = "") -> dict:
|
||
"""List all MLflow experiments with run counts and each latest run's headline IC metrics."""
|
||
store = _open_store(uri)
|
||
|
||
exps = []
|
||
for r in _rows(store, "SELECT experiment_id, name, artifact_location, lifecycle_stage, creation_time, last_update_time FROM experiments ORDER BY creation_time"):
|
||
exp_id = r[0]
|
||
run_count = _rows(store, "SELECT COUNT(*) FROM runs WHERE experiment_id = ?", (exp_id,))[0][0]
|
||
latest = None
|
||
lr = _rows(store, "SELECT run_uuid, name, status, start_time FROM runs WHERE experiment_id = ? ORDER BY start_time DESC LIMIT 1", (exp_id,))
|
||
if lr:
|
||
latest = {
|
||
"run_id": lr[0][0],
|
||
"name": lr[0][1],
|
||
"status": lr[0][2],
|
||
"start_time": lr[0][3],
|
||
"headline": _headline(_latest_metrics_of(store, lr[0][0])),
|
||
}
|
||
exps.append(
|
||
{
|
||
"experiment_id": exp_id,
|
||
"name": r[1],
|
||
"artifact_location": r[2],
|
||
"lifecycle_stage": r[3],
|
||
"creation_time": r[4],
|
||
"last_update_time": r[5],
|
||
"run_count": run_count,
|
||
"latest_run": latest,
|
||
}
|
||
)
|
||
store.close()
|
||
return {"uri": uri, "db": store.db, "experiments": _jsonable(exps)}
|
||
|
||
|
||
def rd_exp_get_experiment(experiment_id: str, uri: str = "") -> dict:
|
||
"""Return one experiment and its runs (meta, params, tags, latest metrics, notes, artifact files)."""
|
||
store = _open_store(uri)
|
||
|
||
exp = _experiment_row(store, experiment_id)
|
||
if exp is None:
|
||
raise ValueError(f"experiment_id {experiment_id!r} not found (see rd_exp_list)")
|
||
runs = []
|
||
for r in _rows(store, "SELECT run_uuid FROM runs WHERE experiment_id = ? ORDER BY start_time DESC", (int(experiment_id),)):
|
||
run_id = r[0]
|
||
meta = _run_meta(store, run_id)
|
||
run_art_root = _artifact_root_for_run(store, int(experiment_id), run_id)
|
||
runs.append(
|
||
{
|
||
"meta": meta,
|
||
"params": _params_of(store, run_id),
|
||
"tags": _tags_of(store, run_id),
|
||
"metrics": _latest_metrics_of(store, run_id),
|
||
"notes": _load_notes(run_art_root, experiment_id, run_id),
|
||
"artifacts": _list_artifact_files(_run_artifact_root(run_art_root, experiment_id, run_id)),
|
||
}
|
||
)
|
||
store.close()
|
||
return {"experiment": _jsonable(exp), "runs": _jsonable(runs)}
|
||
|
||
|
||
def rd_exp_get_run(run_id: str, experiment_id: str = "", uri: str = "") -> dict:
|
||
"""Return one run in full: meta, params, tags, per-step metrics, artifact files and notes."""
|
||
store = _open_store(uri)
|
||
|
||
meta = _run_meta(store, run_id)
|
||
if meta is None:
|
||
raise ValueError(f"run_id {run_id!r} not found")
|
||
exp_id = int(experiment_id) if experiment_id else meta["experiment_id"]
|
||
art_root = _artifact_root_for_run(store, exp_id, run_id)
|
||
out = {
|
||
"meta": meta,
|
||
"params": _params_of(store, run_id),
|
||
"tags": _tags_of(store, run_id),
|
||
"metrics": _step_metrics_of(store, run_id),
|
||
"metrics_latest": _latest_metrics_of(store, run_id),
|
||
"headline": _headline(_latest_metrics_of(store, run_id)),
|
||
"notes": _load_notes(art_root, exp_id, run_id),
|
||
"artifacts": _list_artifact_files(_run_artifact_root(art_root, exp_id, run_id)),
|
||
}
|
||
store.close()
|
||
return _jsonable(out)
|
||
|
||
|
||
def rd_exp_input(run_id: str, experiment_id: str = "", uri: str = "") -> dict:
|
||
"""Return the input configuration of a run: qlib_init, model kwargs, dataset handler kwargs, segments, features, universe. Source is the saved `config` artifact (canonical) or reconstructed."""
|
||
store = _open_store(uri)
|
||
|
||
meta = _run_meta(store, run_id)
|
||
if meta is None:
|
||
raise ValueError(f"run_id {run_id!r} not found")
|
||
exp_id = int(experiment_id) if experiment_id else meta["experiment_id"]
|
||
art_root = _artifact_root_for_run(store, exp_id, run_id)
|
||
store.close()
|
||
out = _run_input(art_root, exp_id, run_id)
|
||
out["run_id"] = run_id
|
||
out["experiment_id"] = exp_id
|
||
return _jsonable(out)
|
||
|
||
|
||
def rd_exp_result(run_id: str, experiment_id: str = "", uri: str = "") -> dict:
|
||
"""Return the results of a run: headline metrics, per-day IC/Rank IC series, group returns, prediction stats, and the backtest report + risk analysis."""
|
||
store = _open_store(uri)
|
||
|
||
meta = _run_meta(store, run_id)
|
||
if meta is None:
|
||
raise ValueError(f"run_id {run_id!r} not found")
|
||
exp_id = int(experiment_id) if experiment_id else meta["experiment_id"]
|
||
art_root = _artifact_root_for_run(store, exp_id, run_id)
|
||
latest = _latest_metrics_of(store, run_id)
|
||
step = _step_metrics_of(store, run_id)
|
||
store.close()
|
||
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
pred = _load_series(root / "pred.pkl")
|
||
label = _load_series(root / "label.pkl") if (root / "label.pkl").exists() else None
|
||
|
||
ic_series, ic_warn = _ic_series(art_root, exp_id, run_id, pred, label)
|
||
grp, grp_warn = _group_returns(art_root, exp_id, run_id, pred, label)
|
||
bt_pts, risk, bt_warn, benchmark, bt_ann = _backtest(art_root, exp_id, run_id)
|
||
pstats, p_warn = _pred_stats(art_root, exp_id, run_id)
|
||
|
||
warnings = [w for w in (ic_warn, grp_warn, bt_warn, p_warn) if w]
|
||
return _jsonable(
|
||
{
|
||
"run_id": run_id,
|
||
"experiment_id": exp_id,
|
||
"metrics": latest,
|
||
"metrics_by_step": step,
|
||
"headline": _headline(latest),
|
||
"ic_series": ic_series,
|
||
"ic_histogram": _ic_histogram(ic_series),
|
||
"monthly_ic": _monthly_ic(ic_series),
|
||
"autocorrelation": _autocorrelation(pred),
|
||
"group_returns": grp,
|
||
"pred_stats": pstats,
|
||
"backtest": {
|
||
"report": bt_pts,
|
||
"risk": risk,
|
||
"benchmark": benchmark,
|
||
"return_annualized": bt_ann.get("return_annualized"),
|
||
"benchmark_annualized": bt_ann.get("benchmark_annualized"),
|
||
},
|
||
"explanation": _explanation(_headline(latest), risk, ic_series),
|
||
"warnings": warnings,
|
||
}
|
||
)
|
||
|
||
|
||
def rd_exp_model(run_id: str, tree_id: int = 0, max_depth: int = DEFAULT_DEPTH, seed: int = 0, experiment_id: str = "", uri: str = "") -> dict:
|
||
"""Return a run's model: hyper-parameters, feature importances, structure stats, and a pruned LightGBM decision tree for a given (seed, tree_id).
|
||
|
||
- ``seed``: for ensemble models, which per-seed sub-model to inspect (0-indexed; default 0). Ignored for plain models.
|
||
- ``tree_id``: which boosting tree to render (default 0). ``tree.milestone_tree_ids`` gives quick jumps (first/mid/best).
|
||
"""
|
||
store = _open_store(uri)
|
||
|
||
meta = _run_meta(store, run_id)
|
||
if meta is None:
|
||
raise ValueError(f"run_id {run_id!r} not found")
|
||
exp_id = int(experiment_id) if experiment_id else meta["experiment_id"]
|
||
art_root = _artifact_root_for_run(store, exp_id, run_id)
|
||
store.close()
|
||
|
||
root = _run_artifact_root(art_root, exp_id, run_id)
|
||
config, _ = _load_pickle(root / "config") if (root / "config").exists() else (None, None)
|
||
hyperparams: Dict[str, Any] = {}
|
||
model_class = None
|
||
model_module = None
|
||
if config:
|
||
model_cfg = (config.get("task", {}) or {}).get("model", {}) or {}
|
||
hyperparams = dict(model_cfg.get("kwargs", {}) or {})
|
||
model_class = model_cfg.get("class")
|
||
model_module = model_cfg.get("module_path")
|
||
|
||
if not hyperparams and not model_class:
|
||
# Legacy / reconstructed runs: recover settings from the fitted model.
|
||
model, _ = _load_model(art_root, exp_id, run_id)
|
||
if model is not None:
|
||
hyperparams = _model_params(model)
|
||
if not model_class:
|
||
model_class = type(model).__name__
|
||
|
||
tree, warn = _model_dump(art_root, exp_id, run_id, int(tree_id), int(max_depth), int(seed))
|
||
out: Dict[str, Any] = {
|
||
"run_id": run_id,
|
||
"experiment_id": exp_id,
|
||
"model_class": model_class,
|
||
"model_module": model_module,
|
||
"hyperparams": _jsonable(hyperparams),
|
||
"tree": tree,
|
||
}
|
||
if warn:
|
||
out["warning"] = warn
|
||
return _jsonable(out)
|
||
|
||
|
||
def rd_exp_blotter(run_id: str, experiment_id: str = "", uri: str = "") -> dict:
|
||
"""Trade-level execution blotter for a run's backtest: FIFO tax-lot realized/unrealized P&L per trade and per position, daily equity/cost curve, position history, and the signal blotter (cross-sectional rank, selected topk, traded flag)."""
|
||
store = _open_store(uri)
|
||
|
||
meta = _run_meta(store, run_id)
|
||
if meta is None:
|
||
raise ValueError(f"run_id {run_id!r} not found")
|
||
exp_id = int(experiment_id) if experiment_id else meta["experiment_id"]
|
||
art_root = _artifact_root_for_run(store, exp_id, run_id)
|
||
store.close()
|
||
|
||
payload, warnings = _blotter(art_root, exp_id, run_id)
|
||
return _jsonable(
|
||
{
|
||
"run_id": run_id,
|
||
"experiment_id": exp_id,
|
||
**payload,
|
||
}
|
||
)
|
||
|
||
|
||
def rd_exp_get_notes(run_id: str, experiment_id: str = "", uri: str = "") -> dict:
|
||
"""Read the hypothesis/evaluation notes recorded for a run (sidecar rd-notes.json)."""
|
||
store = _open_store(uri)
|
||
|
||
meta = _run_meta(store, run_id)
|
||
if meta is None:
|
||
raise ValueError(f"run_id {run_id!r} not found")
|
||
exp_id = int(experiment_id) if experiment_id else meta["experiment_id"]
|
||
art_root = _artifact_root_for_run(store, exp_id, run_id)
|
||
store.close()
|
||
notes = _load_notes(art_root, exp_id, run_id)
|
||
return _jsonable({"run_id": run_id, "experiment_id": exp_id, **notes})
|
||
|
||
|
||
def rd_exp_set_notes(run_id: str, hypothesis: str = "", evaluation: str = "", experiment_id: str = "", uri: str = "") -> dict:
|
||
"""Record hypothesis / evaluation notes for a run (sidecar rd-notes.json next to the run dir)."""
|
||
store = _open_store(uri)
|
||
|
||
meta = _run_meta(store, run_id)
|
||
if meta is None:
|
||
raise ValueError(f"run_id {run_id!r} not found")
|
||
exp_id = int(experiment_id) if experiment_id else meta["experiment_id"]
|
||
art_root = _artifact_root_for_run(store, exp_id, run_id)
|
||
store.close()
|
||
p = _notes_path(art_root, exp_id, run_id)
|
||
p.parent.mkdir(parents=True, exist_ok=True)
|
||
notes = {"hypothesis": hypothesis, "evaluation": evaluation, "updated_at": _dt.datetime.now().isoformat(timespec="seconds")}
|
||
p.write_text(json.dumps(notes, indent=2))
|
||
return _jsonable({"run_id": run_id, "experiment_id": exp_id, "path": str(p), **notes})
|
||
|
||
|
||
# --------------------------------------------------------------------------- deletion
|
||
def _mlflow_store(uri: str):
|
||
"""The real MLflow tracking store (postgres/sqlite) for hard deletes."""
|
||
from mlflow.tracking._tracking_service.utils import _get_store
|
||
|
||
return _get_store(uri)
|
||
|
||
|
||
def _delete_rd_experiments(uri: str, run_ids: List[str]) -> int:
|
||
"""Delete traced ``rd_experiments`` rows referencing the given mlflow run ids.
|
||
|
||
The trace table's FK to ``runs(run_uuid)`` is ``ON DELETE CASCADE``, but we
|
||
delete explicitly so the rows are removed even against a legacy ``SET NULL``
|
||
constraint (or a soft-deleted run that never hits the FK). Returns count.
|
||
"""
|
||
if not run_ids:
|
||
return 0
|
||
if uri.startswith(("postgres://", "postgresql://", "postgresql+psycopg://")):
|
||
import psycopg
|
||
|
||
with psycopg.connect(_postgres_dsn(uri)) as conn, conn.cursor() as cur:
|
||
cur.execute(
|
||
"DELETE FROM rd_experiments WHERE experiment_ref_id = ANY(%s)",
|
||
(run_ids,),
|
||
)
|
||
return cur.rowcount
|
||
db, _art = _sqlite_paths(uri)
|
||
conn = sqlite3.connect(db)
|
||
try:
|
||
ph = ",".join("?" for _ in run_ids)
|
||
cur = conn.execute(
|
||
f"DELETE FROM rd_experiments WHERE experiment_ref_id IN ({ph})", run_ids
|
||
)
|
||
conn.commit()
|
||
return cur.rowcount
|
||
finally:
|
||
conn.close()
|
||
|
||
|
||
def _delete_run_artifacts(art_root: Path, exp_id, run_id: str) -> int:
|
||
"""Remove the on-disk mlflow artifact dir(s) for a run. Returns count removed."""
|
||
removed = 0
|
||
for candidate in (art_root / str(exp_id) / run_id, art_root / run_id):
|
||
if candidate.is_dir():
|
||
shutil.rmtree(candidate, ignore_errors=True)
|
||
removed += 1
|
||
return removed
|
||
|
||
|
||
def _delete_art_root(uri: str) -> Path:
|
||
if uri.startswith(("postgres://", "postgresql://", "postgresql+psycopg://")):
|
||
return _postgres_art_root()
|
||
return _sqlite_paths(uri)[1]
|
||
|
||
|
||
def rd_exp_delete(experiment_id: str, uri: str = "") -> dict:
|
||
"""HARD-delete an MLflow experiment: all its runs + metrics/params/tags, the traced
|
||
``rd_experiments`` rows that reference those runs, and the on-disk artifact files.
|
||
Use ``rd_exp_list`` for ids. Irreversible."""
|
||
from mlflow.entities import ViewType
|
||
|
||
uri = (uri or _default_tracking_uri()).strip()
|
||
store = _mlflow_store(uri)
|
||
try:
|
||
exp = store.get_experiment(experiment_id)
|
||
except Exception as exc: # noqa: BLE001 - not found / bad id
|
||
raise ValueError(f"experiment_id {experiment_id!r} not found (see rd_exp_list): {exc}")
|
||
run_ids = [
|
||
r.info.run_id
|
||
for r in store.search_runs([experiment_id], "", ViewType.ALL, max_results=50000)
|
||
]
|
||
|
||
traced_deleted = _delete_rd_experiments(uri, run_ids)
|
||
art_root = _delete_art_root(uri)
|
||
notes: List[str] = []
|
||
|
||
for run_id in run_ids:
|
||
try:
|
||
store._hard_delete_run(run_id)
|
||
except Exception as exc: # noqa: BLE001 - best-effort per run
|
||
notes.append(f"run {run_id}: {exc}")
|
||
_delete_run_artifacts(art_root, experiment_id, run_id)
|
||
|
||
try:
|
||
store.delete_experiment(experiment_id) # soft (required before hard)
|
||
store._hard_delete_experiment(experiment_id) # purge
|
||
except Exception as exc: # noqa: BLE001 - best-effort
|
||
notes.append(f"experiment: {exc}")
|
||
exp_dir = art_root / str(experiment_id)
|
||
if exp_dir.is_dir():
|
||
shutil.rmtree(exp_dir, ignore_errors=True)
|
||
notes.append(f"removed {exp_dir}")
|
||
|
||
return {
|
||
"experiment_id": experiment_id,
|
||
"experiment_name": getattr(exp, "name", None),
|
||
"runs_deleted": len(run_ids),
|
||
"traced_experiments_deleted": traced_deleted,
|
||
"notes": notes,
|
||
}
|
||
|
||
|
||
def rd_exp_delete_run(run_id: str, experiment_id: str = "", uri: str = "") -> dict:
|
||
"""HARD-delete a single MLflow run (metrics/params/tags), the traced
|
||
``rd_experiments`` row that references it, and its on-disk artifact files.
|
||
Irreversible."""
|
||
uri = (uri or _default_tracking_uri()).strip()
|
||
store = _mlflow_store(uri)
|
||
try:
|
||
run = store.get_run(run_id)
|
||
except Exception as exc: # noqa: BLE001 - not found / bad id
|
||
raise ValueError(f"run_id {run_id!r} not found: {exc}")
|
||
exp_id = experiment_id or str(run.info.experiment_id)
|
||
|
||
traced_deleted = _delete_rd_experiments(uri, [run_id])
|
||
try:
|
||
store._hard_delete_run(run_id)
|
||
except Exception as exc: # noqa: BLE001
|
||
raise ValueError(f"could not hard-delete run {run_id}: {exc}")
|
||
removed = _delete_run_artifacts(_delete_art_root(uri), exp_id, run_id)
|
||
|
||
return {
|
||
"run_id": run_id,
|
||
"experiment_id": exp_id,
|
||
"traced_experiments_deleted": traced_deleted,
|
||
"artifact_dirs_removed": removed,
|
||
}
|
||
|
||
|
||
# --------------------------------------------------------------------------- lineage graph
|
||
def rd_exp_lineage(uri: str = "") -> dict:
|
||
"""Return the full experiment/run lineage as a node+edge graph for the R&D lineage view.
|
||
|
||
Nodes are the traced `rd_experiments` rows (one per evolution step, each
|
||
carrying its own `experiment_ref_id` = an mlflow run) and edges are their
|
||
`evolved_from` self-FK links. Each node is enriched with: the mlflow
|
||
experiment id/name it belongs to, the run's headline metrics (IC/ICIR/
|
||
Rank IC/Rank ICIR from the traced metrics or the run's latest_metrics), its
|
||
git branch / rational / status, and any trading rounds + order counts
|
||
linked to it. This is the same Postgres store as the mlflow tracking db
|
||
when `$DATABASE_URL` is set, so no recursive CTE / PGQ is needed — load all
|
||
nodes and build edges client-side.
|
||
"""
|
||
store = _open_store(uri)
|
||
|
||
nodes: List[dict] = []
|
||
edges: List[dict] = []
|
||
try:
|
||
rows = _rows(store, "SELECT * FROM rd_experiments ORDER BY id")
|
||
cols = [c[0] for c in store.execute("SELECT * FROM rd_experiments LIMIT 0").description or ()]
|
||
except Exception as exc: # noqa: BLE001 - table may not exist yet
|
||
store.close()
|
||
return {"uri": uri, "nodes": [], "edges": [], "note": f"rd_experiments not available: {exc}"}
|
||
|
||
def _col(row: tuple, name: str):
|
||
try:
|
||
return row[cols.index(name)]
|
||
except (ValueError, IndexError):
|
||
return None
|
||
|
||
exp_ids_by_name: Dict[str, Any] = {}
|
||
try:
|
||
for r in _rows(store, "SELECT experiment_id, name FROM experiments"):
|
||
exp_ids_by_name[str(r[1])] = r[0]
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
|
||
for row in rows:
|
||
row_id = _col(row, "id")
|
||
ref_id = _col(row, "experiment_ref_id")
|
||
evolved = _col(row, "evolved_from")
|
||
name = _col(row, "experiment_name")
|
||
metrics = _col(row, "metrics") or {}
|
||
|
||
headline = {}
|
||
if metrics:
|
||
headline = _headline(metrics)
|
||
# Prefer the run's live IC-family metrics when the traced snapshot is
|
||
# a trade-day summary (no IC/ICIR) — the /rd result view reads these.
|
||
if ref_id:
|
||
try:
|
||
run_headline = _headline(_latest_metrics_of(store, ref_id))
|
||
except Exception: # noqa: BLE001
|
||
run_headline = {}
|
||
if run_headline:
|
||
headline = run_headline if not headline else {**headline, **run_headline}
|
||
|
||
# linked trading rounds + order counts (same postgres db)
|
||
rounds = []
|
||
try:
|
||
for r in _rows(
|
||
store,
|
||
"SELECT id, target_date, status FROM trading_rounds WHERE rd_experiment_id = ? ORDER BY target_date",
|
||
(row_id,),
|
||
):
|
||
n_orders = 0
|
||
n_filled = 0
|
||
try:
|
||
cnt = _rows(
|
||
store,
|
||
"SELECT COUNT(*), COUNT(*) FILTER (WHERE status = 'filled') FROM round_orders WHERE round_id = ?",
|
||
(r[0],),
|
||
)[0]
|
||
n_orders, n_filled = cnt[0], cnt[1]
|
||
except Exception: # noqa: BLE001 - orders may be empty
|
||
pass
|
||
rounds.append(
|
||
{
|
||
"round_id": r[0],
|
||
"target_date": str(r[1]) if r[1] is not None else None,
|
||
"status": r[2],
|
||
"orders": n_orders,
|
||
"filled": n_filled,
|
||
}
|
||
)
|
||
except Exception: # noqa: BLE001 - rounds table may not exist
|
||
pass
|
||
|
||
nodes.append(
|
||
{
|
||
"id": row_id,
|
||
"evolved_from": evolved,
|
||
"experiment_name": name,
|
||
"mlflow_experiment_id": exp_ids_by_name.get(str(name)) if name else None,
|
||
"run_id": ref_id,
|
||
"git_branch": _col(row, "git_branch"),
|
||
"rational": _col(row, "rational"),
|
||
"details": _col(row, "details"),
|
||
"evaluation": _col(row, "evaluation"),
|
||
"status": _col(row, "status"),
|
||
"mlruns_dir": _col(row, "mlruns_dir"),
|
||
"start_ts": _col(row, "start_ts"),
|
||
"end_ts": _col(row, "end_ts"),
|
||
"headline": headline,
|
||
"rounds": rounds,
|
||
}
|
||
)
|
||
if evolved:
|
||
edges.append({"source": evolved, "target": row_id})
|
||
|
||
store.close()
|
||
return {"uri": uri, "db": store.db, "nodes": _jsonable(nodes), "edges": _jsonable(edges)}
|
||
|
||
|
||
# --------------------------------------------------------------------------- registration
|
||
def register_tools(server) -> None:
|
||
"""Attach all ``rd_exp_*`` tools to an ``MCPServer`` instance (called by rd_server.py)."""
|
||
for fn in (
|
||
rd_exp_list,
|
||
rd_exp_get_experiment,
|
||
rd_exp_get_run,
|
||
rd_exp_input,
|
||
rd_exp_result,
|
||
rd_exp_model,
|
||
rd_exp_blotter,
|
||
rd_exp_get_notes,
|
||
rd_exp_set_notes,
|
||
rd_exp_delete,
|
||
rd_exp_delete_run,
|
||
rd_exp_lineage,
|
||
):
|
||
server.tool(structured_output=False)(fn)
|