"""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///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 ``/mlruns``. Without it, the legacy unified lake sqlite store ``sqlite:////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: ``/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. ``///artifacts``). Prefer the recorded location when it matches the ``///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 //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)