book: scaffold + ch00 (execution trail as spine) — evidence exp 8-31, round 3

This commit is contained in:
TradeAC Book Agent
2026-08-18 22:35:23 +00:00
commit c93424e76c
83 changed files with 17676 additions and 0 deletions
+18
View File
@@ -0,0 +1,18 @@
"""tac-qlib: read the TradeAC parquet lake from within the qlib research workflow.
This package is fully decoupled from the upstream ``qlib`` checkout. It provides:
- ``tac_qlib.qlib_init.qlib_init``: drop-in ``qlib.init`` configured against the lake.
- ``tac_qlib.data.providers``: qlib data providers (calendar / instruments / features)
backed by the lake parquet files, so ``qlib.init`` + ``D.features`` work without any
``*.bin`` data.
- ``tac_qlib.contrib.data.handler``: a ``DataHandlerLP`` subclass (``TACHandler``) that
builds a train/test dataset from raw OHLCV + pre-computed ta-lib features.
Upstream ``qlib/`` is never modified.
"""
from .qlib_init import qlib_init, provider_config
__all__ = ["qlib_init", "provider_config"]
__version__ = "0.1.0"
File diff suppressed because it is too large Load Diff
+388
View File
@@ -0,0 +1,388 @@
"""Round book MCP server (stdio transport) — the execution trail for algo rounds.
Exposes the round-book tools over MCP so the agent (tac-algo-trade skill) and
the R&D UI can read AND write the same execution trail in Postgres:
scheduler_runs ──► ROUND ──► rd_experiments
round_create / round_update / round_update_status / round_list / round_get windows
fact_record / fact_query evidence
intent_set / intent_get / intent_list target portfolios
decision_record / decision_query / order_link gates + orders
round_sync_fills Alpaca fills
book_reconcile / trail_query / trail_funnel / round_metrics investigation
Run::
.venv/bin/python -m tac_qlib.book_server # stdio MCP server
All tools return JSON-safe dicts. Logging goes to stderr; stdout is reserved
for the MCP protocol. DB access is via tac_qlib.book_db (psycopg + DATABASE_URL);
fill sync additionally needs APCA_API_KEY_ID / APCA_API_SECRET_KEY when the
agent does not pass `orders` explicitly.
"""
from __future__ import annotations
import functools
import json
import sys
from typing import Any, Dict, List, Optional, Sequence
from mcp.server.mcpserver import MCPServer
from tac_qlib import book_db
book_db._load_repo_env()
server = MCPServer(
name="tac-rd-book",
title="TradeAC round book",
instructions=(
"Execution trail for algo trading rounds on the TradeAC stack: create "
"round windows, record evidence facts, set versioned target intents, "
"record placed/skipped decisions, link Alpaca orders, sync fills, and "
"reconcile / trace the funnel. Backed by Postgres (DATABASE_URL)."
),
version="0.1.0",
)
def _log(message: str) -> None:
print(f"[tac-rd-book] {message}", file=sys.stderr)
def _as_obj(value: Any) -> Any:
"""Accept structured MCP input as JSON strings or as already-parsed dicts/lists."""
if isinstance(value, str):
if not value.strip():
return None
try:
return json.loads(value)
except json.JSONDecodeError:
return value
return value
def _obj(value: Any) -> Optional[Dict[str, Any]]:
parsed = _as_obj(value)
return parsed if isinstance(parsed, dict) else None
def _arr(value: Any) -> Optional[List[Any]]:
parsed = _as_obj(value)
return parsed if isinstance(parsed, list) else None
def _open_round(func):
"""Ensure the round tables exist before any round-book operation.
Uses ``functools.wraps`` so ``inspect.signature`` follows ``__wrapped__``
and the MCP tool schema keeps the real typed parameters."""
@functools.wraps(func)
def wrapper(*args, **kwargs):
try:
book_db.ensure_schema()
except Exception as exc: # noqa: BLE001
_log(f"schema check failed: {exc}")
return func(*args, **kwargs)
return wrapper
# --------------------------------------------------------------------------- windows
@_open_round
def round_create(
target_date: str,
source: str = "scheduled",
signal_date: str = "",
scheduler_run_id: int = 0,
rd_experiment_id: int = 0,
experiment_name: str = "",
run_id: str = "",
model_path: str = "",
strategy_snapshot: str = "{}",
account_equity_at_sizing: float = 0.0,
) -> dict:
"""Open a round window for a target trading date. Idempotent per
(source, target_date): an already-open round for the same window is
returned unchanged (``reused=True``). Returns the full round row."""
snap = _obj(strategy_snapshot) or {}
return book_db.create_round(
source=source,
target_date=target_date,
signal_date=signal_date or None,
scheduler_run_id=scheduler_run_id or None,
rd_experiment_id=rd_experiment_id or None,
experiment_name=experiment_name or None,
run_id=run_id or None,
model_path=model_path or None,
strategy_snapshot=snap,
account_equity_at_sizing=account_equity_at_sizing or None,
)
def round_update(
round_id: int,
source: str = "",
target_date: str = "",
signal_date: str = "",
scheduler_run_id: int = 0,
rd_experiment_id: int = 0,
experiment_name: str = "",
run_id: str = "",
model_path: str = "",
strategy_snapshot: str = "",
account_equity_at_sizing: float = 0.0,
) -> dict:
"""Update a round window's metadata — e.g. pin the new training run
(``run_id`` / ``model_path``) and strategy snapshot after the retrain.
Empty / zero values leave the field unchanged."""
return book_db.update_round(
round_id,
source=source or None,
target_date=target_date or None,
signal_date=signal_date or None,
scheduler_run_id=scheduler_run_id or None,
rd_experiment_id=rd_experiment_id or None,
experiment_name=experiment_name or None,
run_id=run_id or None,
model_path=model_path or None,
strategy_snapshot=_obj(strategy_snapshot) if strategy_snapshot else None,
account_equity_at_sizing=account_equity_at_sizing or None,
)
def round_update_status(
round_id: int,
status: str = "",
locked_intent_id: int = 0,
summary_metrics: str = "{}",
feedback_note: str = "",
) -> dict:
"""Advance a round (open → locked → settled | aborted). ``locked_intent_id``
pins the intent reconciliation uses. ``summary_metrics`` / ``feedback_note``
update the round summary."""
return book_db.update_round_status(
round_id,
status=status or None,
locked_intent_id=locked_intent_id or None,
summary_metrics=_obj(summary_metrics),
feedback_note=feedback_note or None,
)
def round_list(
source: str = "",
target_date: str = "",
status: str = "",
limit: int = 20,
include_detail: bool = False,
) -> dict:
"""List round windows (newest first), optionally filtered by source /
target_date / status. ``include_detail`` attaches each round's funnel
counts + roll-up metrics (used by the /dashboard/rounds list)."""
return {
"rounds": book_db.list_rounds(
source=source or None,
target_date=target_date or None,
status=status or None,
limit=limit,
with_detail=bool(include_detail),
)
}
def round_get(round_id: int) -> dict:
"""Full detail of one round: window row + intents, decisions, orders, facts,
funnel and reconciliation — everything the UI's round detail page needs."""
round_row = book_db.get_round(round_id)
if round_row is None:
return {"error": f"round {round_id} not found"}
return {
"round": round_row,
"intents": book_db.list_intents(round_id),
"decisions": book_db.query_decisions(round_id),
"orders": book_db.list_orders(round_id),
"facts": book_db.query_facts(round_id, limit=500),
"funnel": book_db.funnel(round_id),
"reconcile": book_db.reconcile(round_id),
"metrics": book_db.metrics(round_id),
}
# --------------------------------------------------------------------------- facts
def fact_record(round_id: int, kind: str, payload: str = "{}", symbol: str = "", source: str = "") -> dict:
"""Append an evidence event (signal_score, quote, news_sentiment,
account_state, position_state, strategy_config, ...) to a round."""
return book_db.record_fact(
round_id,
kind=kind,
payload=_obj(payload),
symbol=symbol or None,
source=source or None,
)
def fact_query(round_id: int, kind: str = "", symbol: str = "", limit: int = 200) -> dict:
"""Query a round's recorded facts (newest first), optionally filtered by kind/symbol."""
return {"facts": book_db.query_facts(round_id, kind=kind or None, symbol=symbol or None, limit=limit)}
# --------------------------------------------------------------------------- intents
def intent_set(round_id: int, target_portfolio: str, raw_strategy_output: str = "{}", reason: str = "") -> dict:
"""Write the next target-portfolio version for a round (auto-increments and
supersedes the previous active version). ``target_portfolio`` is a JSON
array of {symbol, side, qty, notional, expected_price, score, rank, weight}."""
return book_db.set_intent(
round_id,
target_portfolio=_arr(target_portfolio) or [],
raw_strategy_output=_obj(raw_strategy_output),
reason=reason or None,
)
def intent_get(round_id: int, version: int = 0) -> dict:
"""Get a round's intent — the given version, or the active (max) version
when ``version`` is omitted."""
intent = book_db.get_intent(round_id, version=version or None)
return {"intent": intent} if intent else {"intent": None, "error": f"no intent for round {round_id}"}
def intent_list(round_id: int) -> dict:
"""List every target-portfolio version for a round (oldest first)."""
return {"intents": book_db.list_intents(round_id)}
# --------------------------------------------------------------------------- decisions / orders
def decision_record(
round_id: int,
symbol: str,
side: str,
qty: float = 0.0,
order_type: str = "",
expected_price: float = 0.0,
status: str = "intended",
reason: str = "",
reason_detail: str = "",
intent_id: int = 0,
supersedes_decision_id: int = 0,
alpaca_order_id: str = "",
client_order_id: str = "",
) -> dict:
"""Record one per-symbol decision by the gates. Use status ``skipped`` /
``rejected`` with a ``reason`` for deliberate skips; placed orders carry
``alpaca_order_id`` / ``client_order_id`` (an execution row is created).
Passing ``supersedes_decision_id`` marks the previous decision superseded."""
return book_db.record_decision(
round_id,
symbol=symbol,
side=side,
qty=qty or None,
order_type=order_type or None,
expected_price=expected_price or None,
status=status,
reason=reason or None,
reason_detail=reason_detail or None,
intent_id=intent_id or None,
supersedes_decision_id=supersedes_decision_id or None,
alpaca_order_id=alpaca_order_id or None,
client_order_id=client_order_id or None,
)
def decision_query(round_id: int, symbol: str = "", include_superseded: bool = True) -> dict:
"""List a round's decisions, optionally filtered by symbol."""
return {"decisions": book_db.query_decisions(round_id, symbol=symbol or None, include_superseded=include_superseded)}
def order_link(
round_id: int,
decision_id: int,
alpaca_order_id: str = "",
client_order_id: str = "",
qty_filled: float = -1.0,
avg_fill_price: float = -1.0,
status: str = "",
) -> dict:
"""Create or update the execution row for a placed decision (idempotent per
decision). Use ``qty_filled=-1`` to leave the value unchanged."""
return book_db.link_order(
round_id,
decision_id=decision_id,
alpaca_order_id=alpaca_order_id or None,
client_order_id=client_order_id or None,
qty_filled=qty_filled if qty_filled >= 0 else None,
avg_fill_price=avg_fill_price if avg_fill_price >= 0 else None,
status=status or None,
)
def round_sync_fills(round_id: int, orders: str = "", feed: str = "iex") -> dict:
"""Pull Alpaca order state into the round. Pass ``orders`` as a JSON array
(tac-engine ``list_orders`` output) or omit it to fetch from Alpaca with
APCA_API_* env vars. Marks orders superseded when the effective intent no
longer targets their symbol+side. Returns updated / superseded / unmatched."""
parsed = _arr(orders) if isinstance(orders, str) else orders
return book_db.sync_fills(round_id, orders=parsed if isinstance(parsed, list) else None, feed=feed or "iex")
# --------------------------------------------------------------------------- investigation
def book_reconcile(round_id: int) -> dict:
"""Reconcile the round: effective intent targets vs decisions vs fills, with
per-symbol residuals and roll-ups (cash/BP impact, slippage bps, cost)."""
return book_db.reconcile(round_id)
def trail_query(round_id: int, symbol: str = "") -> dict:
"""Per-symbol waterfall: intent target → decision → order → fill."""
return {"trail": book_db.trail(round_id, symbol=symbol or None)}
def trail_funnel(round_id: int) -> dict:
"""Decision funnel counts for a round: targets → decided → placed → filled,
plus skipped-reason breakdown and superseded count."""
return book_db.funnel(round_id)
def round_metrics(round_id: int) -> dict:
"""Roll-up metrics: placed/filled order counts, invested notional, turnover."""
return book_db.metrics(round_id)
def register_tools(mcp_server: MCPServer) -> None:
"""Attach all round-book tools to an ``MCPServer`` instance."""
for fn in (
round_create,
round_update,
round_update_status,
round_list,
round_get,
fact_record,
fact_query,
intent_set,
intent_get,
intent_list,
decision_record,
decision_query,
order_link,
round_sync_fills,
book_reconcile,
trail_query,
trail_funnel,
round_metrics,
):
mcp_server.tool(structured_output=False)(fn)
def main() -> None:
register_tools(server)
server.run()
if __name__ == "__main__":
main()
+11
View File
@@ -0,0 +1,11 @@
from . import data # noqa: F401 (registers tac_qlib.contrib.data)
from . import model, strategy # noqa: F401
from .data import TACHandler # noqa: F401
from .model import RankICLGBModel # noqa: F401
from .strategy import OptimalStopControl # noqa: F401
__all__ = [
"TACHandler",
"RankICLGBModel",
"OptimalStopControl",
]
@@ -0,0 +1,3 @@
from .handler import TACHandler
__all__ = ["TACHandler"]
+254
View File
@@ -0,0 +1,254 @@
"""TACHandler: a qlib DataHandlerLP that builds datasets from the TradeAC lake.
This is the "custom DataHandler" entry point (Option B): the handler is referenced from the
workflow yaml's ``dataset.handler`` and reads OHLCV + pre-computed ta-lib features straight
from the lake parquet files through ``QLibDataLoader`` + the tac_qlib feature provider.
The standard qlib processor pipeline (``infer_processors`` / ``learn_processors``) still runs
on top, so existing recipes such as ``DropnaLabel``, ``CSZScoreNorm`` or ``RobustZScoreNorm``
keep working unchanged.
"""
from __future__ import annotations
import os
from inspect import getfullargspec
from typing import List, Optional, Tuple, Union
from qlib.data.dataset import processor as processor_module
from qlib.data.dataset.handler import DataHandlerLP
from qlib.utils import get_callable_kwargs
from ...data.config import (
LakeConfig,
timeframe_for_freq,
NON_FEATURE_COLUMNS,
)
DEFAULT_INFER_PROCESSORS = [
{"class": "DropAllNaN", "kwargs": {}},
{"class": "ProcessInf", "kwargs": {}},
{"class": "ZScoreNorm", "kwargs": {}},
{"class": "Fillna", "kwargs": {}},
]
DEFAULT_LEARN_PROCESSORS = [
{"class": "DropnaLabel"},
{"class": "CSZScoreNorm", "kwargs": {"fields_group": "label"}},
]
#: always include raw OHLCV; ta-lib columns are discovered from the lake and appended.
RAW_FEATURE_FIELDS = ("$open", "$high", "$low", "$close", "$vwap", "$volume")
DEFAULT_LABEL = "Ref($close,-2)/Ref($close,-1)-1"
def check_transform_proc(proc_l, fit_start_time, fit_end_time):
"""Port of ``qlib.contrib.data.handler.check_transform_proc`` (inject fit window into procs)."""
new_l = []
for p in proc_l:
if not isinstance(p, processor_module.Processor):
klass, pkwargs = get_callable_kwargs(p, processor_module)
args = getfullargspec(klass).args
if "fit_start_time" in args and "fit_end_time" in args:
assert fit_start_time is not None and fit_end_time is not None, (
"Make sure `fit_start_time` and `fit_end_time` are not None."
)
pkwargs.update({"fit_start_time": fit_start_time, "fit_end_time": fit_end_time})
proc_config = {"class": klass.__name__, "kwargs": pkwargs}
if isinstance(p, dict) and "module_path" in p:
proc_config["module_path"] = p["module_path"]
new_l.append(proc_config)
else:
new_l.append(p)
return new_l
def get_common_feature_fields(lake_root=None, market="US", timeframe="1d") -> List[str]:
"""Discover feature columns present in *every* feature file of the lake.
Walks the `family=ta|sp` partition layout (plus any legacy flat files).
TA and SP columns are disjoint by construction, so the common set is
computed per family (columns shared by all symbol files of that family),
then the per-family results are unioned. Returns sorted field names
(without the ``$`` prefix). Empty if no features are persisted.
"""
cfg = LakeConfig(lake_root, market)
feat_dir = cfg.features_dir(timeframe)
if not feat_dir.exists():
return []
import pyarrow.parquet as pq
def _family_common(fam_dir: Path) -> set:
common = None
for p in sorted(fam_dir.glob("symbol=*.parquet")):
try:
cols = set(pq.read_schema(p).names) - set(NON_FEATURE_COLUMNS)
except Exception: # pragma: no cover - skip unreadable files
continue
common = cols if common is None else (common & cols)
if not common:
break
return common or set()
common: set = set()
# family tier: features/market=*/timeframe=*/family=*/symbol=*.parquet
for fam in ("ta", "sp"):
fam_dir = feat_dir / f"family={fam}"
if fam_dir.is_dir():
common |= _family_common(fam_dir)
# legacy flat: features/market=*/timeframe=*/symbol=*.parquet
if (feat_dir / "family=ta").exists() or (feat_dir / "family=sp").exists():
pass # family layout already covered
else:
common |= _family_common(feat_dir)
return sorted(common)
class DropAllNaN(processor_module.Processor):
"""Drop feature columns that are all-NaN over the fit window.
The lake can hold fully-empty indicator columns (e.g. a ta-lib output that was NaN
from the start). Such columns carry no learnable signal and make ``ZScoreNorm.fit``
warn on empty slices, so we drop them before any other processor runs. The drop set
is fixed on the fit window once (during ``fit``), then applied consistently to every
segment so train/valid/test keep identical feature columns.
"""
def __init__(self, fit_start_time=None, fit_end_time=None):
self.fit_start_time = fit_start_time
self.fit_end_time = fit_end_time
self.cols_to_drop = []
def fit(self, df=None):
if df is None or len(df) == 0:
return self
window = df
if self.fit_start_time is not None and self.fit_end_time is not None:
try:
from qlib.data.dataset.utils import fetch_df_by_index
window = fetch_df_by_index(
df, slice(self.fit_start_time, self.fit_end_time), level="datetime"
)
except Exception: # pragma: no cover - defensive
window = df
if len(window) == 0:
return self
self.cols_to_drop = [c for c in window.columns if window[c].isna().all()]
return self
def __call__(self, df):
if self.cols_to_drop:
return df.drop(columns=self.cols_to_drop, errors="ignore")
return df
class TACHandler(DataHandlerLP):
"""DataHandlerLP backed by the TradeAC parquet lake.
Parameters mirror ``Alpha158``: ``instruments``/``start_time``/``end_time``/``freq`` define
the queried window; ``feature_fields`` selects the features (default: raw OHLCV + all common
ta-lib columns found in the lake); ``label`` is a qlib expression for the target.
"""
def __init__(
self,
instruments="all",
start_time=None,
end_time=None,
freq="day",
infer_processors=DEFAULT_INFER_PROCESSORS,
learn_processors=DEFAULT_LEARN_PROCESSORS,
fit_start_time=None,
fit_end_time=None,
process_type=DataHandlerLP.PTYPE_A,
filter_pipe=None,
feature_fields=None,
label=DEFAULT_LABEL,
lake_root=None,
market="US",
**kwargs,
):
# default the processor fit window to the queried window (like Alpha158 without a split)
if fit_start_time is None:
fit_start_time = start_time
if fit_end_time is None:
fit_end_time = end_time
infer_processors = check_transform_proc(infer_processors, fit_start_time, fit_end_time)
learn_processors = check_transform_proc(learn_processors, fit_start_time, fit_end_time)
feature_fields = self._normalize_feature_fields(feature_fields, freq, lake_root, market)
if not feature_fields:
raise ValueError(
"no feature fields available for the lake; set `feature_fields` explicitly "
"(e.g. ['$close', '$rsi_14', '$sma_20'])"
)
label_expr, label_names = self._normalize_label(label)
data_loader = {
"class": "QlibDataLoader",
"kwargs": {
"config": {
"feature": (feature_fields, feature_fields),
"label": (label_expr, label_names),
},
"filter_pipe": filter_pipe,
"freq": freq,
},
}
super().__init__(
instruments=instruments,
start_time=start_time,
end_time=end_time,
data_loader=data_loader,
infer_processors=infer_processors,
learn_processors=learn_processors,
process_type=process_type,
**kwargs,
)
# ------------------------------------------------------------------ config
@staticmethod
def _normalize_feature_fields(feature_fields, freq, lake_root, market) -> List[str]:
if feature_fields is None:
common = get_common_feature_fields(lake_root, market, timeframe_for_freq(freq))
feature_fields = list(RAW_FEATURE_FIELDS) + ["$" + f for f in common if "$" + f not in RAW_FEATURE_FIELDS]
elif isinstance(feature_fields, str):
feature_fields = [f.strip() for f in feature_fields.split(",") if f.strip()]
fields = [f if f.startswith("$") else "$" + f for f in feature_fields]
# de-dup while preserving order
seen, out = set(), []
for f in fields:
if f not in seen:
seen.add(f)
out.append(f)
return out
@staticmethod
def _normalize_label(label) -> Tuple[List[str], List[str]]:
if isinstance(label, str):
return [label], ["LABEL0"]
if isinstance(label, (list, tuple)):
if len(label) == 2 and isinstance(label[0], str):
return [label[0]], list(label[1]) if isinstance(label[1], (list, tuple)) else [label[1]]
return list(label), ["LABEL%d" % i for i in range(len(label))]
raise TypeError(f"unsupported label config: {label!r}")
# ------------------------------------------------------------------ utils
def get_label_config(self):
return DEFAULT_LABEL
@staticmethod
def discover_feature_fields(lake_root=None, market="US", freq="day") -> List[str]:
return get_common_feature_fields(lake_root, market, timeframe_for_freq(freq))
__all__ = ["TACHandler", "DropAllNaN", "get_common_feature_fields"]
# Make `DropAllNaN` resolvable by bare name from processor configs (e.g. the default
# ``infer_processors`` and workflow yamls that reference it without a ``module_path``),
# mirroring how qlib registers its own processors in ``qlib.data.dataset.processor``.
processor_module.DropAllNaN = DropAllNaN
@@ -0,0 +1,4 @@
from .rank_ensemble import RankICEnsembleLGBModel # noqa: F401
from .rank_gbdt import RankICLGBModel, rankic_feval # noqa: F401
__all__ = ["RankICLGBModel", "rankic_feval", "RankICEnsembleLGBModel"]
@@ -0,0 +1,189 @@
"""Seed-ensembled LightGBM that early-stops on cross-sectional RankIC.
``RankICEnsembleLGBModel`` wraps ``RankICLGBModel`` (per-day RankIC feval +
``metric='None'`` + ``first_metric_only`` early stopping) over a seed ensemble:
one sub-model is trained per seed with identical hyper-parameters, and
predictions are averaged across seeds. This is the model class the
``tac-rd-rank-ensemble-isolated`` reference run wires into its workflow
(``module_path: tac_qlib.contrib.model.rank_ensemble``).
The ensemble inherits the RankIC early-stopping behaviour of the single-seed
model (valid RankIC drives the stopping iteration) while the seed averaging
stabilizes the prediction against any single seed's early-stopping path.
Training is parallelized: the seed sub-models train in a thread pool —
``lgb.train`` is C++ and releases the GIL, so concurrent seeds do not block on
the GIL (5 seeds ~40min/5 on this box). Measured on a 6-physical-core / 12 SMT
host: the seeds scale ~2x, not linearly — the runs are memory-bandwidth bound
and each Booster caps its threads at ``cores // workers`` so 5 concurrent
boosters don't oversubscribe; larger-core hosts scale better. The qlib data
pipeline is warmed once on the calling thread (fills the handler cache), and
each worker then prepares its **own** ``lgb.Dataset`` (independent handle, so
no concurrent ``construct()`` on a shared handle — LightGBM's ``Dataset`` is
not thread-safe to build). qlib's ``R`` recorder is also not thread-safe, so
the per-seed evaluation curves are logged on the calling thread after the pool
finishes.
Wired into a workflow yaml like:
model:
class: RankICEnsembleLGBModel
module_path: tac_qlib.contrib.model.rank_ensemble
kwargs:
loss: mse
learning_rate: 0.02
num_leaves: 31
n_estimators: 3000
num_boost_round: 3000
early_stopping_rounds: 200
min_data_in_leaf: 20
lambda_l2: 0.5
colsample_bytree: 0.8
subsample: 0.8
subsample_freq: 1
reg_alpha: 0.1
reg_lambda: 1.0
seeds: "42,7,2026,99,123"
parallel: 5
Any ``**kwargs`` other than ``seeds``/``parallel`` are forwarded unchanged to
every ``RankICLGBModel`` sub-model (same params, different ``seed``).
"""
from __future__ import annotations
import os
from concurrent.futures import ThreadPoolExecutor
from typing import List, Optional
import pandas as pd
from qlib.data.dataset import DatasetH
from qlib.data.dataset.handler import DataHandlerLP
from tac_qlib.contrib.model.rank_gbdt import RankICLGBModel
__all__ = ["RankICEnsembleLGBModel"]
class RankICEnsembleLGBModel(RankICLGBModel):
"""Seed ensemble of RankIC-early-stopping LightGBM models.
Parameters
----------
seeds : comma-separated integers, one sub-model per seed.
parallel : number of seeds to train concurrently. ``0`` (default) = auto
(all seeds, bounded by the available cores); ``1`` = sequential.
**kwargs : forwarded to every ``RankICLGBModel`` sub-model (model
hyper-parameters). ``seeds``/``parallel`` are consumed here and not
forwarded.
"""
def __init__(self, seeds: str = "42", parallel: int = 0, **kwargs):
self.seeds = [int(s.strip()) for s in str(seeds).split(",") if s.strip()]
if not self.seeds:
raise ValueError("seeds must contain at least one integer")
self.parallel = int(parallel)
# drop seed/parallel handling from the base kwargs, keep everything else
self._model_kwargs = dict(kwargs)
super().__init__(**self._model_kwargs)
self._models: List[RankICLGBModel] = []
# --------------------------------------------------------------- helpers
@staticmethod
def _cores() -> int:
try:
return max(1, len(os.sched_getaffinity(0)))
except AttributeError:
return max(1, os.cpu_count() or 1)
def _worker_count(self) -> int:
if self.parallel > 0:
return min(len(self.seeds), self.parallel)
return min(len(self.seeds), self._cores())
# ------------------------------------------------------------------ fit
def fit(
self,
dataset: DatasetH,
num_boost_round: Optional[int] = None,
early_stopping_rounds: Optional[int] = None,
verbose_eval: int = 20,
evals_result=None,
reweighter=None,
**kwargs,
):
"""Train one RankICLGBModel per seed and keep them for prediction.
The qlib data pipeline is warmed once on this thread (handler cache),
then each seed sub-model trains in a parallel worker thread on its own
``lgb.Dataset`` (LightGBM releases the GIL in ``lgb.train``). Evals
are logged on this thread after the pool (qlib's ``R`` is not
thread-safe).
"""
n_round = num_boost_round or self.num_boost_round
n_es = early_stopping_rounds or self.early_stopping_rounds
if len(self.seeds) == 1:
m = RankICLGBModel(seed=self.seeds[0], **self._model_kwargs)
m.fit(
dataset,
num_boost_round=n_round,
early_stopping_rounds=n_es,
verbose_eval=verbose_eval,
evals_result=evals_result,
reweighter=reweighter,
**kwargs,
)
self._models = [m]
return
# Warm the qlib handler cache once on this thread so the workers'
# concurrent prepare() calls only hit cached frames (no first-write race).
proto = RankICLGBModel(seed=self.seeds[0], **self._model_kwargs)
proto._prepare_data(dataset, reweighter)
workers = self._worker_count()
# Cap per-Booster threads so concurrent seeds don't oversubscribe
# (LightGBM's num_threads=0 uses ALL cores per Booster).
per_booster = max(1, self._cores() // workers)
def fit_seed(seed):
m = RankICLGBModel(seed=seed, **self._model_kwargs)
if workers > 1 and "num_threads" not in m.params:
m.params["num_threads"] = per_booster
ds_l = m._prepare_data(dataset, reweighter)
booster, evals, names = m._train_from_datasets(
ds_l,
num_boost_round=n_round,
early_stopping_rounds=n_es,
verbose_eval=verbose_eval,
**kwargs,
)
m.model = booster
return m, evals, names
with ThreadPoolExecutor(max_workers=workers) as ex:
results = list(ex.map(fit_seed, self.seeds))
self._models = [m for m, _, _ in results]
# Merge + log evals on the main thread (qlib's R is not thread-safe).
if evals_result is not None:
for m, evals, names in results:
for k in names:
for key, val in evals.get(k, {}).items():
evals_result.setdefault(f"{k}.seed{m.params['seed']}", {})[key] = val
for m, evals, names in results:
self._log_evals(evals, names, prefix=f"seed{m.params['seed']}.")
# -------------------------------------------------------------- predict
def predict(self, dataset: DatasetH, segment="test") -> pd.Series:
"""Average the per-seed predictions over the given segment."""
if not self._models:
raise ValueError("model is not fitted yet!")
preds = [m.predict(dataset, segment=segment) for m in self._models]
if len(preds) == 1:
return preds[0]
frame = pd.concat(preds, axis=1)
return frame.mean(axis=1)
@@ -0,0 +1,200 @@
"""LGBModel variant that early-stops on cross-sectional RankIC instead of l2.
Standard qlib ``LGBModel`` early-stops on the regression loss (mse). For
cross-sectional alpha signals the quantity we actually care about is the per-day
rank correlation (Rank IC), which mse early-stopping does not optimize for.
Experiments on the 50-ETF lake (SP-5d 55-feature panel) show that early-stopping
on a custom RankIC feval lifts RankIC 0.047 -> 0.075 vs. the mse-stopped model.
This class reuses ``LGBModel``'s data preparation but:
- tags each ``lgb.Dataset`` with per-day query ``group`` sizes so a ranking
metric can be computed per trading day;
- injects a custom ``feval`` (mean per-day Spearman of pred vs label) into
``lgb.train``; early stopping then selects the iteration that maximizes
RankIC on the valid set;
- forces ``metric='None'`` + ``first_metric_only=True`` so early-stopping
tracks RankIC only (not the regression loss).
Wired into a workflow yaml like:
model:
class: RankICLGBModel
module_path: tac_qlib.contrib.model.rank_gbdt
kwargs:
loss: mse
learning_rate: 0.03
num_leaves: 31
n_estimators: 500
...
The rank feval is used for early-stopping selection only; the objective stays
the configured loss (default mse). Set ``rank_eval=False`` to fall back to the
plain LGBModel behaviour (early-stop on the loss).
Generic: works for any cross-sectional panel whose qlib dataset index has a
``datetime`` level (each level value = one query group). The per-day groups are
derived automatically, so no universe-specific configuration is needed.
"""
from __future__ import annotations
from typing import List, Optional, Tuple
import numpy as np
import pandas as pd
import lightgbm as lgb
from qlib.data.dataset import DatasetH
from qlib.data.dataset.handler import DataHandlerLP
from qlib.contrib.model.gbdt import LGBModel
from qlib.workflow import R
__all__ = ["RankICLGBModel", "rankic_feval"]
def _per_day_spearman(preds: np.ndarray, labels: np.ndarray, group: np.ndarray) -> float:
"""Mean per-day Spearman rank correlation of preds vs labels.
``group`` holds the number of rows of each trading day (query group), in
order. Days with <3 valid rows or a constant pred/label are skipped.
"""
if group is None or len(group) == 0:
return 0.0
offs = np.concatenate([[0], np.cumsum(group.astype(int))])
vals = []
for i in range(len(group)):
s = slice(offs[i], offs[i + 1])
p, l = preds[s], labels[s]
if len(p) < 3 or np.std(p) == 0 or np.std(l) == 0:
continue
vals.append(np.corrcoef(pd.Series(p).rank(), pd.Series(l).rank())[0, 1])
return float(np.mean(vals)) if vals else 0.0
def rankic_feval(preds, dataset):
"""LightGBM feval: mean RankIC (higher is better in lgb convention)."""
labels = dataset.get_label()
group = dataset.get_group()
ric = _per_day_spearman(preds, labels, group)
return "rankic", ric, True # (name, value, higher_is_better)
class RankICLGBModel(LGBModel):
"""LGBModel that early-stops on per-day RankIC via a custom feval."""
def __init__(self, rank_eval: bool = True, **kwargs):
super().__init__(**kwargs)
self.rank_eval = rank_eval
def _prepare_data(self, dataset: DatasetH, reweighter=None) -> List[Tuple[lgb.Dataset, str]]:
ds_l = []
assert "train" in dataset.segments
for key in ["train", "valid"]:
if key in dataset.segments:
df = dataset.prepare(key, col_set=["feature", "label"], data_key=DataHandlerLP.DK_L)
if df.empty:
raise ValueError("Empty data from dataset, please check your dataset config.")
x, y = df["feature"], df["label"]
if y.values.ndim == 2 and y.values.shape[1] == 1:
y = np.squeeze(y.values)
else:
raise ValueError("LightGBM doesn't support multi-label training")
if reweighter is None:
w = None
elif hasattr(reweighter, "reweight"):
w = reweighter.reweight(df)
else:
raise ValueError("Unsupported reweighter type.")
# per-day query groups: each trading day is one group
if self.rank_eval and isinstance(df.index, pd.MultiIndex) and "datetime" in df.index.names:
group = df.groupby(level="datetime").size().to_numpy(dtype=np.int32)
else:
group = None
d = lgb.Dataset(x.values, label=y, weight=w, group=group, free_raw_data=False)
ds_l.append((d, key))
return ds_l
def _train_from_datasets(
self,
ds_l: List[Tuple[lgb.Dataset, str]],
num_boost_round: Optional[int] = None,
early_stopping_rounds: Optional[int] = None,
verbose_eval: int = 20,
evals_result=None,
**kwargs,
) -> Tuple[lgb.Booster, dict, List[str]]:
"""Train a Booster from already-prepared ``lgb.Dataset`` objects.
Pure training — no ``R.log_metrics`` — so it can be called from worker
threads (qlib's ``R`` recorder is not thread-safe; the caller decides
when/where to log). Returns ``(booster, evals_result, segment_names)``.
"""
if evals_result is None:
evals_result = {}
ds, names = list(zip(*ds_l))
callbacks = [
lgb.early_stopping(
self.early_stopping_rounds if early_stopping_rounds is None else early_stopping_rounds
),
lgb.log_evaluation(period=verbose_eval),
lgb.record_evaluation(evals_result),
]
if self.rank_eval:
# early-stopping must be driven ONLY by the RankIC feval, not l2.
# metric='None' suppresses the default l2 metric; first_metric_only
# makes early_stopping track the single remaining (rankic) metric.
self.params["metric"] = "None"
self.params["first_metric_only"] = True
feval = rankic_feval
else:
self.params.pop("metric", None)
self.params.pop("first_metric_only", None)
feval = None
booster = lgb.train(
self.params,
ds[0],
num_boost_round=self.num_boost_round if num_boost_round is None else num_boost_round,
valid_sets=ds,
valid_names=names,
feval=feval,
callbacks=callbacks,
**kwargs,
)
return booster, evals_result, list(names)
def _log_evals(self, evals_result, names: List[str], prefix: str = "") -> None:
"""Log recorded evaluation curves to qlib's active recorder."""
for k in names:
for key, val in evals_result.get(k, {}).items():
name = f"{prefix}{key}.{k}"
for epoch, m in enumerate(val):
R.log_metrics(**{name.replace("@", "_"): m}, step=epoch)
def fit(
self,
dataset: DatasetH,
num_boost_round: Optional[int] = None,
early_stopping_rounds: Optional[int] = None,
verbose_eval: int = 20,
evals_result=None,
reweighter=None,
**kwargs,
):
if evals_result is None:
evals_result = {}
ds_l = self._prepare_data(dataset, reweighter)
self.model, evals_result, names = self._train_from_datasets(
ds_l,
num_boost_round=num_boost_round,
early_stopping_rounds=early_stopping_rounds,
verbose_eval=verbose_eval,
evals_result=evals_result,
**kwargs,
)
self._log_evals(evals_result, names)
@@ -0,0 +1,3 @@
from .optimal_stop import OptimalStopControl # noqa: F401
__all__ = ["OptimalStopControl"]
@@ -0,0 +1,217 @@
"""Optimal-stopping / stochastic-control strategy for cross-sectional signals.
Entry is a control policy: a symbol opens a position only when its cross-sectional
signal percentile is at or above ``entry_pct`` (i.e. it is one of the top-ranked
names) and the portfolio has fewer than ``topk`` open positions.
Exit is an optimal-stopping rule: a held position is stopped (closed) when its
signal percentile falls below ``exit_pct`` (the continuation value of holding is
no longer worth the risk), OR after ``max_hold_days`` (time stop / finite
horizon), OR when the position P&L breaches ``sl`` (loss control) and the
position has been held at least ``min_hold_days``.
Sizing is fixed ``notional`` per position (equal-weight control), unlike the
TopkDropout cash-allocation heuristic.
Wired into qrun workflows like any ``BaseStrategy`` (see ``PortAnaRecord``
config). Mirrors the API usage of qlib's ``TopkDropoutStrategy``: ``Order``/
``OrderDir`` from ``qlib.backtest.decision``, ``trade_calendar`` /
``trade_exchange`` / ``trade_position`` injected by the backtest executor.
"""
from __future__ import annotations
from typing import List
import pandas as pd
from qlib.backtest import Order
from qlib.backtest.decision import OrderDir, TradeDecisionWO
from qlib.contrib.strategy.signal_strategy import BaseSignalStrategy
__all__ = ["OptimalStopControl"]
DEFAULT_NOTIONAL = 20_000.0
DEFAULT_ENTRY_PCT = 0.80
DEFAULT_EXIT_PCT = 0.50
DEFAULT_MAX_HOLD_DAYS = 10
DEFAULT_MIN_HOLD_DAYS = 2
DEFAULT_SL = -0.06
class OptimalStopControl(BaseSignalStrategy):
"""Optimal-stopping long-only strategy over a cross-sectional signal.
Parameters
----------
topk : max number of concurrent positions.
entry_pct : min cross-sectional score percentile required to OPEN (0..1).
exit_pct : held positions are stopped when score percentile < exit_pct.
max_hold_days : hard time stop (finite-horizon close).
min_hold_days : minimum holding days before stop-loss is evaluated.
notional : $ per position (equal-weight control).
sl : stop-loss threshold as fraction of entry price (<= 0), disabled if 0.
"""
def __init__(
self,
*,
signal=None,
topk: int = 10,
entry_pct: float = DEFAULT_ENTRY_PCT,
exit_pct: float = DEFAULT_EXIT_PCT,
max_hold_days: int = DEFAULT_MAX_HOLD_DAYS,
min_hold_days: int = DEFAULT_MIN_HOLD_DAYS,
notional: float = DEFAULT_NOTIONAL,
sl: float = DEFAULT_SL,
risk_degree: float = 0.95,
trade_exchange=None,
level_infra=None,
common_infra=None,
**kwargs,
):
super().__init__(
signal=signal,
trade_exchange=trade_exchange,
level_infra=level_infra,
common_infra=common_infra,
**kwargs,
)
self.topk = topk
self.entry_pct = entry_pct
self.exit_pct = exit_pct
self.max_hold_days = max_hold_days
self.min_hold_days = min_hold_days
self.notional = notional
self.sl = sl
# ------------------------------------------------------------------ utils
@staticmethod
def _pct_rank(score: pd.Series) -> pd.Series:
return score.rank(pct=True)
def _entry_price(self, pos) -> float:
# Position stores avg entry price under key "price" (see Position.position)
price = pos.position.get("price")
if price is None:
price = pos.get_stock_amount("price")
return float(price)
def _pnl_pct(self, pos, mark: float) -> float:
entry = self._entry_price(pos)
if not entry or entry != entry:
return 0.0
return mark / entry - 1.0
def _is_tradable(self, code, start, end, direction) -> bool:
try:
return self.trade_exchange.is_stock_tradable(
stock_id=code, start_time=start, end_time=end, direction=direction
)
except TypeError: # some exchanges take no direction kwarg
return self.trade_exchange.is_stock_tradable(stock_id=code, start_time=start, end_time=end)
# ------------------------------------------------------------ decision
def generate_trade_decision(self, execute_result=None):
trade_step = self.trade_calendar.get_trade_step()
trade_start, trade_end = self.trade_calendar.get_step_time(trade_step)
pred_start, pred_end = self.trade_calendar.get_step_time(trade_step, shift=1)
pred_score = self.signal.get_signal(start_time=pred_start, end_time=pred_end)
if isinstance(pred_score, pd.DataFrame):
pred_score = pred_score.iloc[:, 0]
if pred_score is None or len(pred_score) == 0:
return TradeDecisionWO([], self)
pct = self._pct_rank(pred_score)
time_per_step = self.trade_calendar.get_freq()
current_temp = __import__("copy").deepcopy(self.trade_position)
holdings = {}
for code in current_temp.get_stock_list():
if abs(current_temp.get_stock_amount(code)) > 1e-6:
holdings[code] = current_temp
# ---- optimal stopping: close held positions -----------------------
sell_orders: List[Order] = []
closed_today = set()
kept = {}
for code, pos in holdings.items():
held = current_temp.get_stock_count(code, bar=time_per_step)
mark = self.trade_exchange.get_deal_price(
stock_id=code, start_time=trade_start, end_time=trade_end, direction=Order.SELL
)
if mark is None or mark != mark:
continue
rank = pct.get(code, 0.0)
stop_pnl = held >= self.min_hold_days and self.sl < 0 and self._pnl_pct(pos, mark) <= self.sl
if held >= self.max_hold_days or rank < self.exit_pct or stop_pnl:
amt = abs(current_temp.get_stock_amount(code))
o = Order(stock_id=code, amount=amt, start_time=trade_start,
end_time=trade_end, direction=Order.SELL)
if self.trade_exchange.check_order(o):
sell_orders.append(o)
self.trade_exchange.deal_order(o, position=current_temp)
closed_today.add(code)
else:
kept[code] = mark
# ---- equal-weight control: target notional per name -----------------
# candidate opens: top-ranked names whose signal pct >= entry_pct
rank_desc = pred_score.sort_values(ascending=False)
held_codes = set(kept)
opens = []
for sym in rank_desc.index:
if len(opens) >= self.topk:
break
if sym in held_codes:
continue
if pct.get(sym, 0.0) < self.entry_pct:
continue
if not self._is_tradable(sym, trade_start, trade_end, OrderDir.BUY):
continue
opens.append(sym)
targets = held_codes | set(opens)
if not targets:
return TradeDecisionWO(sell_orders, self)
# total value (cash + marked positions) -> per-target notional
total_value = current_temp.get_cash()
for code, mark in kept.items():
total_value += abs(current_temp.get_stock_amount(code)) * mark
target_notional = total_value * self.risk_degree / max(1, len(targets))
# ---- rebalance kept positions toward target weight ------------------
buy_orders: List[Order] = []
for code, mark in kept.items():
cur = abs(current_temp.get_stock_amount(code)) * mark
diff_notional = target_notional - cur
if abs(diff_notional) / target_notional < 0.02:
continue # skip tiny rebalances
amount_delta = diff_notional / mark
direction = Order.BUY if amount_delta > 0 else Order.SELL
o = Order(stock_id=code, amount=abs(amount_delta), start_time=trade_start,
end_time=trade_end, direction=direction)
if self.trade_exchange.check_order(o):
(buy_orders if direction == Order.BUY else sell_orders).append(o)
self.trade_exchange.deal_order(o, position=current_temp)
# ---- open new positions at target weight ----------------------------
for sym in opens:
px = self.trade_exchange.get_deal_price(
stock_id=sym, start_time=trade_start, end_time=trade_end, direction=OrderDir.BUY
)
if px is None or px != px or px <= 0:
continue
amount = target_notional / px
factor = self.trade_exchange.get_factor(
stock_id=sym, start_time=trade_start, end_time=trade_end
)
amount = self.trade_exchange.round_amount_by_trade_unit(amount, factor)
o = Order(stock_id=sym, amount=amount, start_time=trade_start,
end_time=trade_end, direction=Order.BUY)
if self.trade_exchange.check_order(o):
buy_orders.append(o)
return TradeDecisionWO(sell_orders + buy_orders, self)
+25
View File
@@ -0,0 +1,25 @@
from .config import (
LakeConfig,
BAR_FIELD_MAP,
FREQ_TO_TIMEFRAME,
UNKNOWN_FIELD_NAMES,
timeframe_for_freq,
resolve_lake_root,
)
from .providers import (
LakeCalendarProvider,
LakeInstrumentProvider,
LakeFeatureProvider,
)
__all__ = [
"LakeConfig",
"BAR_FIELD_MAP",
"FREQ_TO_TIMEFRAME",
"UNKNOWN_FIELD_NAMES",
"timeframe_for_freq",
"resolve_lake_root",
"LakeCalendarProvider",
"LakeInstrumentProvider",
"LakeFeatureProvider",
]
+202
View File
@@ -0,0 +1,202 @@
"""TradeAC lake configuration helpers.
The lake is a hive-partitioned parquet store (see ``tac-engine/skills/tradeac-lake``):
$TAC_LAKE_DIR/
├── market=US/
│ └── timeframe=1d/
│ └── symbol=AAPL.parquet # OHLCV bars: t, date, o, h, l, c, v, n, vw
├── features/ # indicators, wide format, family tier
│ └── market=US/
│ └── timeframe=1d/
│ ├── family=ta/symbol=AAPL.parquet # t, sma_5, sma_20, rsi_14, ...
│ └── family=sp/symbol=AAPL.parquet # t, sp_ou_*, sp_hmm_*, ...
├── calendar.parquet # trading days per market
├── coverage.parquet # per (market,timeframe,symbol) loaded windows
└── symbols.parquet # asset master
"""
from __future__ import annotations
import os
from pathlib import Path
from typing import Dict, List, Optional
import pandas as pd
#: qlib freq string (Freq.__str__) -> lake timeframe partition name
FREQ_TO_TIMEFRAME: Dict[str, str] = {
"day": "1d",
"1d": "1d",
"min": "1m",
"1min": "1m",
"5min": "5m",
"10min": "10m",
"15min": "15m",
"30min": "30m",
"hour": "1h",
"1hour": "1h",
"2hour": "2h",
"4hour": "4h",
"week": "1w",
"1week": "1w",
"month": "1M",
"1month": "1M",
}
#: bar-field map: qlib field name (without the leading ``$``) -> lake bar column
BAR_FIELD_MAP: Dict[str, str] = {
"open": "o",
"high": "h",
"low": "l",
"close": "c",
"volume": "v",
"vwap": "vw",
"avg_amount": "vw", # amount / volume
}
#: fields that qlib core/backtest queries but the lake does not store -> all-NaN
UNKNOWN_FIELD_NAMES = ("factor", "change", "trade_unit", "suspend_flag")
#: columns in the parquet files that are not features
NON_FEATURE_COLUMNS = ("t", "date", "market", "timeframe", "symbol")
def timeframe_for_freq(freq: str) -> str:
"""Map a qlib frequency (e.g. ``day``, ``1min``) to a lake timeframe (e.g. ``1d``)."""
f = str(freq).lower()
if f not in FREQ_TO_TIMEFRAME:
raise ValueError(
f"unsupported qlib freq {freq!r}; supported freqs: {sorted(set(FREQ_TO_TIMEFRAME))}"
)
return FREQ_TO_TIMEFRAME[f]
def resolve_lake_root(lake_root: Optional[str] = None) -> Path:
"""Resolve the lake root: explicit arg > ``TAC_LAKE_DIR`` (no fallback).
``TAC_LAKE_DIR`` is **mandatory** — there is deliberately no default
A missing/empty value raises so a
misconfigured environment never silently points at a wrong directory.
"""
if lake_root is None:
lake_root = os.environ.get("TAC_LAKE_DIR")
if not lake_root:
raise RuntimeError(
"TAC_LAKE_DIR is not set. Point it at the TradeAC lake root, e.g. "
"export TAC_LAKE_DIR=/home/data/lake (docker) or set an absolute "
"path in your local .env."
)
return Path(str(lake_root)).expanduser().resolve()
class LakeConfig:
"""Path helpers + cached readers for a (lake_root, market) combination."""
def __init__(self, lake_root: Optional[str] = None, market: str = "US"):
self.lake_root: Path = resolve_lake_root(lake_root)
self.market: str = (market or "US").upper()
# ---- paths --------------------------------------------------------------
def bar_dir(self, timeframe: str) -> Path:
return self.lake_root / f"market={self.market}" / f"timeframe={timeframe}"
def bar_path(self, timeframe: str, symbol: str) -> Path:
return self.bar_dir(timeframe) / f"symbol={str(symbol).upper()}.parquet"
def features_dir(self, timeframe: str) -> Path:
return self.lake_root / "features" / f"market={self.market}" / f"timeframe={timeframe}"
def features_path(self, timeframe: str, symbol: str) -> Path:
# Legacy flat path (no family tier). Prefer `load_features` which
# resolves the family=ta|sp partition layout.
return self.features_dir(timeframe) / f"symbol={str(symbol).upper()}.parquet"
def load_features(self, timeframe: str, symbol: str) -> pd.DataFrame:
"""All feature columns for a symbol, merging the `family=ta` and
`family=sp` partitions by timestamp. Returns an empty frame when no
feature files exist (legacy flat layout falls back transparently)."""
sym = str(symbol).upper()
frames = []
for family in ("ta", "sp"):
p = self.features_dir(timeframe) / f"family={family}" / f"symbol={sym}.parquet"
if p.exists():
frames.append(pd.read_parquet(p))
if not frames:
flat = self.features_dir(timeframe) / f"symbol={sym}.parquet"
if flat.exists():
return pd.read_parquet(flat)
return pd.DataFrame()
if len(frames) == 1:
return frames[0]
merged = frames[0]
for extra in frames[1:]:
merged = merged.merge(extra, on="t", how="outer", suffixes=("", "_dup"))
for c in [c for c in merged.columns if c.endswith("_dup")]:
merged = merged.drop(columns=c)
return merged
def calendar_path(self) -> Path:
return self.lake_root / "calendar.parquet"
def symbols_path(self) -> Path:
return self.lake_root / "symbols.parquet"
def coverage_path(self) -> Path:
return self.lake_root / "coverage.parquet"
# ---- metadata readers ----------------------------------------------------
def load_symbols(self) -> List[str]:
"""All symbols known to the lake (from ``symbols.parquet``)."""
p = self.symbols_path()
if not p.exists():
return []
df = pd.read_parquet(p)
if "symbol" not in df.columns:
return []
return sorted(df["symbol"].astype(str).str.upper().tolist())
def symbol_spans(self, symbol: str, timeframe: str) -> List[tuple]:
"""Listing span(s) ``[(start_iso, end_iso)]`` for a symbol from coverage.parquet."""
p = self.coverage_path()
if p.exists():
try:
df = pd.read_parquet(p)
except Exception: # pragma: no cover - defensive
df = pd.DataFrame()
if len(df):
df = df[
(df.get("market") == self.market)
& (df.get("timeframe") == timeframe)
& (df.get("symbol") == str(symbol).upper())
]
if len(df):
row = df.iloc[0]
first = pd.Timestamp(row["first_t"]).date()
last = pd.Timestamp(row["last_t"]).date()
return [(first.isoformat(), last.isoformat())]
# fallback: derive from the bar file itself
p = self.bar_path(timeframe, symbol)
if p.exists():
import pyarrow.parquet as pq
tbl = pq.read_table(p, columns=["t"])
first = pd.Timestamp(tbl.column("t")[0].as_py()).date()
last = pd.Timestamp(tbl.column("t")[-1].as_py()).date()
return [(first.isoformat(), last.isoformat())]
return [("1970-01-01", "2099-12-31")]
def load_calendar_dates(self) -> List[pd.Timestamp]:
"""Trading days (midnight timestamps) for the market, from ``calendar.parquet``."""
p = self.calendar_path()
if p.exists():
df = pd.read_parquet(p)
if "date" in df.columns:
if "market" in df.columns:
df = df[df["market"] == self.market]
dates = pd.to_datetime(df["date"]).dt.normalize().sort_values().unique()
return [pd.Timestamp(x) for x in dates]
return []
def __repr__(self) -> str: # pragma: no cover
return f"LakeConfig(lake_root={self.lake_root}, market={self.market})"
+230
View File
@@ -0,0 +1,230 @@
"""qlib data providers backed by the TradeAC parquet lake.
These providers plug into the standard qlib mechanism: ``qlib.init(calendar_provider=...,
instrument_provider=..., feature_provider=...)`` instantiates them and binds them to the
``Cal`` / ``Inst`` / ``FeatureD`` wrappers (see ``qlib.data.data.register_all_wrappers``).
The rest of qlib (``LocalDatasetProvider`` expression engine, backtest ``Exchange``) keeps
working unchanged because the interface contract is identical to the file-based providers:
- ``feature()`` returns a ``pd.Series`` indexed by the **calendar position** range
``[start_index, end_index]`` (matching ``FileFeatureStorage.__getitem__`` semantics).
- ``list_instruments()`` returns ``{symbol: [(start, end), ...]}``.
- ``load_calendar()`` returns a list of ``pd.Timestamp`` trading days.
"""
from __future__ import annotations
import bisect
from typing import Dict, List, Optional, Union
import numpy as np
import pandas as pd
from qlib.data.data import CalendarProvider, FeatureProvider, InstrumentProvider
from qlib.log import get_module_logger
from .config import (
BAR_FIELD_MAP,
LakeConfig,
UNKNOWN_FIELD_NAMES,
timeframe_for_freq,
)
logger = get_module_logger("tac_qlib.data.providers")
def _day_freq(freq: str) -> bool:
return str(freq).lower() in ("day", "1d")
def _calendar_keys(cal: List[pd.Timestamp], freq: str) -> pd.Index:
"""Convert calendar timestamps into the same key space as the lake parquet."""
if _day_freq(freq):
return pd.Index([pd.Timestamp(x).date() for x in cal])
return pd.Index([pd.Timestamp(x) for x in cal])
class LakeCalendarProvider(CalendarProvider):
"""Trading calendar read from ``<lake>/calendar.parquet`` (fallback: derived from bars)."""
def __init__(self, lake_root: Optional[str] = None, market: str = "US"):
super().__init__()
self.cfg = LakeConfig(lake_root, market)
def load_calendar(self, freq, future):
timeframe = timeframe_for_freq(freq)
if not _day_freq(freq):
raise NotImplementedError(
f"freq={freq!r} (timeframe={timeframe}) is not supported yet: the lake calendar "
f"only covers daily sessions; add a minute-level calendar to `calendar.parquet`"
)
dates = self.cfg.load_calendar_dates()
if not dates:
# Fallback: derive the trading-day set from the persisted bar files.
bar_dir = self.cfg.bar_dir(timeframe)
if bar_dir.exists():
import pyarrow.parquet as pq
cal: Dict[pd.Timestamp, None] = {}
for p in sorted(bar_dir.glob("symbol=*.parquet")):
tbl = pq.read_table(p, columns=["t"])
for v in tbl.column("t"):
cal[pd.Timestamp(v.as_py()).normalize()] = None
dates = sorted(cal.keys())
if not dates:
return []
if future:
# append the next calendar day so that "today" is a valid trade date
last = dates[-1]
dates = dates + [pd.Timestamp(last) + pd.Timedelta(days=1)]
return dates
class LakeInstrumentProvider(InstrumentProvider):
"""Instruments from ``<lake>/symbols.parquet`` with listing spans from ``coverage.parquet``."""
def __init__(
self,
lake_root: Optional[str] = None,
market: str = "US",
markets: Optional[Dict[str, list]] = None,
):
super().__init__()
self.cfg = LakeConfig(lake_root, market)
#: optional named pools, e.g. ``{"sp500": ["AAPL", "MSFT"], "etf": ["SPY"]}``.
#: ``all`` / any unregistered name resolves to every symbol in the lake.
self.markets: Dict[str, list] = markets or {}
def _resolve_symbols(self, market: Union[str, list]) -> List[str]:
if isinstance(market, (list, tuple, pd.Index, np.ndarray)):
return [str(s).upper() for s in market]
if isinstance(market, str) and "," in market:
return [s.strip().upper() for s in market.split(",") if s.strip()]
if market in self.markets:
return [str(s).upper() for s in self.markets[market]]
return self.cfg.load_symbols()
def list_instruments(self, instruments, start_time=None, end_time=None, freq="day", as_list=False):
market = instruments["market"]
timeframe = timeframe_for_freq(freq)
symbols = self._resolve_symbols(market)
if not symbols:
if as_list:
return []
return {}
# clip listing spans to the queried window (mirror of LocalInstrumentProvider)
from qlib.data.data import Cal # pylint: disable=C0415
cal = Cal.calendar(freq=freq)
start_time = pd.Timestamp(start_time or cal[0])
end_time = pd.Timestamp(end_time or cal[-1])
out: Dict[str, list] = {}
for symbol in symbols:
spans = []
for begin, end in self.cfg.symbol_spans(symbol, timeframe):
lo = max(start_time, pd.Timestamp(begin))
hi = min(end_time, pd.Timestamp(end))
if lo <= hi:
spans.append((lo, hi))
if spans:
out[symbol] = spans
filter_pipe = instruments.get("filter_pipe") or []
for filter_config in filter_pipe:
from qlib.data import filter as F # pylint: disable=C0415
filter_t = getattr(F, filter_config["filter_type"]).from_config(filter_config)
out = filter_t(out, start_time, end_time, freq)
if as_list:
return list(out)
return out
class LakeFeatureProvider(FeatureProvider):
"""Feature data from the lake parquet (OHLCV bars + pre-computed ta-lib features).
Field routing:
- ``$open/$high/$low/$close/$volume/$vwap`` -> bar parquet columns
- ``$amount`` (= v*vw), ``$avg_amount`` (= vw) -> derived from bar parquet
- ``$factor/$change/...`` -> all-NaN (not stored)
- anything else -> a ta-lib column in the features parquet
"""
def __init__(self, lake_root: Optional[str] = None, market: str = "US"):
super().__init__()
self.cfg = LakeConfig(lake_root, market)
self._bar_cache: Dict[tuple, pd.DataFrame] = {}
self._feature_cache: Dict[tuple, pd.DataFrame] = {}
# ------------------------------------------------------------------ caches
def _load_bar_df(self, instrument: str, timeframe: str) -> pd.DataFrame:
key = (instrument, timeframe)
if key not in self._bar_cache:
p = self.cfg.bar_path(timeframe, instrument)
self._bar_cache[key] = pd.read_parquet(p) if p.exists() else pd.DataFrame()
return self._bar_cache[key]
def _load_feature_df(self, instrument: str, timeframe: str) -> pd.DataFrame:
key = (instrument, timeframe)
if key not in self._feature_cache:
self._feature_cache[key] = self.cfg.load_features(timeframe, instrument)
return self._feature_cache[key]
@staticmethod
def _keys(df: pd.DataFrame, freq: str) -> pd.Index:
ts = pd.to_datetime(df["t"])
return ts.dt.date if _day_freq(freq) else ts
# ------------------------------------------------------------------ fields
def _extract(self, instrument: str, field: str, timeframe: str, freq: str) -> Optional[pd.Series]:
"""Return the field as a Series keyed by date/timestamp (None if not present in the lake)."""
bar = self._load_bar_df(instrument, timeframe)
if field in BAR_FIELD_MAP:
col = BAR_FIELD_MAP[field]
if col in bar.columns:
return bar[col].astype(float).set_axis(self._keys(bar, freq))
return None
if field == "amount":
if "v" in bar.columns and "vw" in bar.columns:
return (bar["v"] * bar["vw"]).astype(float).set_axis(self._keys(bar, freq))
return None
if field in UNKNOWN_FIELD_NAMES:
return None
feat = self._load_feature_df(instrument, timeframe)
if field in feat.columns:
return feat[field].astype(float).set_axis(self._keys(feat, freq))
return None
# ------------------------------------------------------------------ api
def _get_calendar(self, freq: str) -> List[pd.Timestamp]:
from qlib.data.data import Cal # pylint: disable=C0415
cal = Cal.calendar(freq=freq)
return list(cal)
def feature(self, instrument, field, start_index, end_index, freq):
field = str(field)[1:]
timeframe = timeframe_for_freq(freq)
cal = self._get_calendar(freq)
n = len(cal)
lo = max(0, int(start_index))
hi = min(n - 1, int(end_index))
if lo > hi:
return pd.Series(dtype=np.float32)
keys = _calendar_keys(cal[lo : hi + 1], freq)
ser = self._extract(str(instrument).upper(), field, timeframe, freq)
if ser is None:
vals = np.full(len(keys), np.nan, dtype=np.float64)
else:
vals = ser.reindex(keys).to_numpy(dtype=np.float64)
return pd.Series(vals, index=pd.RangeIndex(lo, hi + 1))
+64
View File
@@ -0,0 +1,64 @@
"""Drop-in replacement for ``qlib.init`` that configures qlib against the TradeAC lake.
Usage::
from tac_qlib.qlib_init import qlib_init
qlib_init(
provider_uri="/path/to/lake", # same layout as tac-engine's TAC_LAKE_DIR
market="US",
freq="day",
markets={"sp500": ["AAPL", "MSFT"]}, # optional named instrument pools
**qlib_init_kwargs, # anything qlib.init accepts
)
It sets ``provider_uri`` to the lake root and points the calendar / instrument / feature
providers at the lake-backed implementations, then delegates to the upstream ``qlib.init``.
The dataset provider (expression engine, backtest ``Exchange``) is left untouched, so the
rest of the qlib workflow is byte-for-byte upstream code.
"""
from __future__ import annotations
from typing import Dict, List, Optional, Union
import qlib
from qlib.config import C
from .data.config import LakeConfig, resolve_lake_root
PROVIDERS = "tac_qlib.data.providers"
def provider_config(cls: str, **kwargs) -> dict:
return {"class": f"{PROVIDERS}.{cls}", "kwargs": kwargs}
def qlib_init(
provider_uri: Optional[str] = None,
market: str = "US",
freq: str = "day",
markets: Optional[Dict[str, list]] = None,
**qlib_kwargs,
) -> qlib.Initialized:
"""Initialize qlib with the lake-backed data providers and re-export qlib.init results."""
if qlib_kwargs.pop("calendar_provider", None) is not None or qlib_kwargs.pop("instrument_provider", None) is not None:
raise ValueError("calendar_provider / instrument_provider are managed by tac_qlib; use `market` instead")
lake_root = resolve_lake_root(provider_uri)
if freq != "day":
raise ValueError("freq must be 'day' for now: the lake calendar only covers daily sessions")
qlib_kwargs.setdefault("provider_uri", lake_root)
qlib_kwargs.setdefault("region", "us")
qlib_kwargs.setdefault("expression_cache", None)
qlib_kwargs.setdefault("dataset_cache", None)
qlib_kwargs["calendar_provider"] = provider_config("LakeCalendarProvider", lake_root=lake_root, market=market)
qlib_kwargs["instrument_provider"] = provider_config(
"LakeInstrumentProvider", lake_root=lake_root, market=market, markets=markets or {}
)
qlib_kwargs["feature_provider"] = provider_config("LakeFeatureProvider", lake_root=lake_root, market=market)
return qlib.init(**qlib_kwargs)
__all__ = ["qlib_init", "provider_config"]
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+144
View File
@@ -0,0 +1,144 @@
"""Risk-limit spec shared by backtest and live executor.
One JSON spec is consulted by BOTH ``rd_backtest`` (as a strategy filter
overlay) and ``rd_strategy_targets`` (as pre-gate + sizing caps), so a limit
that holds in backtest holds in live — the round's ``strategy_snapshot``
stores the exact spec used.
Supported keys (all optional, all pct are 0-100):
liquidity_floor_adv : min avg daily dollar volume (USD) per symbol.
Names below it are filtered out of the tradable set.
size_cap_pct : max notional per name as % of account equity.
concentration_cap_pct: max total deployed as % of account equity.
drawdown_pause_pct : if equity drawdown from peak exceeds this, new buys
are paused (executor gate; not expressible in a
one-shot qlib backtest and therefore documented).
"""
from __future__ import annotations
import json
import os
from typing import Any, Dict, List, Optional, Tuple
import pandas as pd
def parse_limits(spec: Optional[str]) -> Dict[str, float]:
"""Parse a risk_limits JSON string into a flat float map (empty = no limits)."""
if not spec or not str(spec).strip():
return {}
if isinstance(spec, dict):
raw = spec
else:
raw = json.loads(str(spec))
out: Dict[str, float] = {}
for k in ("liquidity_floor_adv", "size_cap_pct", "concentration_cap_pct", "drawdown_pause_pct"):
v = raw.get(k)
if v is not None and str(v) != "":
out[k] = float(v)
return out
def dollar_adv(
symbols: List[str],
lake_root: str = "",
market: str = "US",
asof: Optional[str] = None,
lookback: int = 20,
) -> Dict[str, float]:
"""Average daily dollar volume per symbol over the ``lookback`` sessions
ending at ``asof`` (inclusive), read straight from lake 1d bars. Symbols
with no lake data map to 0.0 (treated as illiquid)."""
from tac_qlib.data.config import LakeConfig, resolve_lake_root
cfg = LakeConfig(resolve_lake_root(lake_root or None), market)
asof_ts = pd.Timestamp(asof) if asof else pd.Timestamp.utcnow()
out: Dict[str, float] = {}
for sym in sorted({str(s).upper() for s in symbols}):
p = cfg.bar_path("1d", sym)
if not p.exists():
out[sym] = 0.0
continue
try:
df = pd.read_parquet(p)
except Exception:
out[sym] = 0.0
continue
if not len(df):
out[sym] = 0.0
continue
tcol = df["t"] if "t" in df.columns else df["date"]
ts = pd.to_datetime(tcol)
df = df.assign(_t=ts).sort_values("_t")
df = df[df["_t"] <= asof_ts]
if not len(df):
out[sym] = 0.0
continue
df = df.tail(lookback)
px = df["vw"] if "vw" in df.columns else df["c"]
out[sym] = float((df["v"] * px).mean()) if len(df) else 0.0
return out
def apply_to_ranking(
ranking: pd.Series,
adv: Dict[str, float],
limits: Dict[str, float],
account: float,
risk_degree: float,
topk: int,
) -> Tuple[pd.Series, Dict[str, Any]]:
"""Executor-side overlay on the ranked signal (``pd.Series`` symbol -> score).
Returns (filtered_ranking, applied) where ``filtered_ranking`` has
illiquid names removed and ``applied`` records what the limits did (audit
trail). Per-name notional and total caps are reported but not folded into
the ranking — the caller sizes targets and can read ``applied`` to cap.
"""
applied: Dict[str, Any] = {"notes": [], "dropped_liquidity": []}
filtered = ranking
floor = limits.get("liquidity_floor_adv")
if floor:
dropped = [s for s in filtered.index if adv.get(str(s).upper(), 0.0) < floor]
if dropped:
filtered = filtered.drop(index=[s for s in dropped if s in filtered.index])
applied["dropped_liquidity"] = [str(s) for s in dropped]
applied["notes"].append(f"liquidity floor ${floor:,.0f} ADV dropped {len(dropped)}")
per_name = account * risk_degree / max(topk, 1)
size_cap = limits.get("size_cap_pct")
if size_cap:
cap = account * size_cap / 100.0
applied["size_cap_notional"] = round(cap, 2)
if per_name > cap:
applied["per_name_capped_from"] = round(per_name, 2)
per_name = cap
applied["notes"].append(f"size cap {size_cap:g}% cut per-name notional to ${cap:,.2f}")
applied["per_name_notional"] = round(per_name, 2)
n_buys = min(topk, max(len(filtered), 0))
conc = limits.get("concentration_cap_pct")
if conc:
conc_cap = account * conc / 100.0
applied["concentration_cap_notional"] = round(conc_cap, 2)
total = per_name * max(n_buys, 1)
if total > conc_cap:
applied["total_capped_from"] = round(total, 2)
applied["notes"].append(f"concentration cap {conc:g}% cut total to ${conc_cap:,.2f}")
per_name = conc_cap / max(n_buys, 1)
applied["per_name_capped_from"] = applied.get("per_name_capped_from") or round(total / max(n_buys, 1), 2)
applied["per_name_notional"] = round(per_name, 2)
applied["total_notional"] = round(min(total, conc_cap), 2)
else:
applied["total_notional"] = round(per_name * n_buys, 2)
return filtered, applied
def drawdown_pause(equity: float, peak_equity: float, limits: Dict[str, float]) -> Tuple[bool, Optional[str]]:
"""Executor gate: True when drawdown from peak exceeds drawdown_pause_pct."""
pct = limits.get("drawdown_pause_pct")
if not pct or not peak_equity or not equity:
return False, None
dd = (peak_equity - equity) / peak_equity * 100.0
if dd >= pct:
return True, f"drawdown {dd:.1f}% >= pause {pct:g}% (peak ${peak_equity:,.2f}, equity ${equity:,.2f})"
return False, None
+635
View File
@@ -0,0 +1,635 @@
"""Trace tools for the tac-qlib-rd MCP server — replace the `trace.sh`/`trace_db.py`/`git_exp.sh` scripts.
The R&D lineage (`/rd/lineage`) and the round book build on the `rd_experiments`
Postgres table. This module exposes the full trace lifecycle as MCP tools so an
agent can drive tracing through the long-lived tac-qlib-rd server instead of
shelling out to bash scripts (which re-import psycopg + reconnect per call and
force the agent to parse prose output).
Because the server is long-lived, `psycopg` is imported and the DB connection
is opened once per call (not once per script invocation), and every tool returns
a single JSON object — no output parsing, fully deterministic.
Git operations (fork / commit / push on the `experiments/` clone) are performed
via `git` subprocess with the repo's mandated credential helper, exactly as the
old `git_exp.sh` did.
Env (from the repo `.env`, already loaded by rd_server): DATABASE_URL,
EMBEDDING_API_BASE_URL, EMBEDDING_API_KEY, GIT_USER, GIT_PASS, GIT_REPO_URL,
TAC_LAKE_DIR.
"""
from __future__ import annotations
import json
import os
import sqlite3
import subprocess
import sys
from pathlib import Path
from typing import Any, Dict, List, Optional
import psycopg
from psycopg.rows import dict_row
from tac_qlib.trace_embed import embed
EMBEDDING_DIM = 384
_MIN_SCORE = 0.5
# --------------------------------------------------------------------------- db
def _conn():
url = (os.environ.get("DATABASE_URL") or "").strip()
if not url:
raise RuntimeError("DATABASE_URL is not set")
return psycopg.connect(url, row_factory=dict_row)
def _now_iso() -> str:
from datetime import datetime, timezone
return datetime.now(timezone.utc).isoformat()
def _vector_literal(vec: Optional[List[float]]) -> Optional[str]:
if not vec:
return None
return "[" + ",".join(repr(float(v)) for v in vec) + "]"
def _jsonable(obj: Any) -> Any:
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 hasattr(obj, "isoformat"):
return obj.isoformat()
return obj
def _row_json(row: Dict[str, Any]) -> Dict[str, Any]:
out: Dict[str, Any] = {}
for k, v in row.items():
if k in ("rational_embedding", "details_embedding"):
out[k] = v.tolist() if hasattr(v, "tolist") else v
elif k == "metrics" and isinstance(v, str):
try:
out[k] = json.loads(v)
except Exception: # noqa: BLE001
out[k] = v
else:
out[k] = v
return out
def _get_row(exp_id: int) -> Optional[Dict[str, Any]]:
with _conn() as conn, conn.cursor() as cur:
cur.execute("SELECT * FROM rd_experiments WHERE id = %s", (exp_id,))
return cur.fetchone()
def _all_rows(limit: int) -> List[Dict[str, Any]]:
with _conn() as conn, conn.cursor() as cur:
cur.execute("SELECT * FROM rd_experiments ORDER BY id DESC LIMIT %s", (limit,))
return cur.fetchall()
def _init_db() -> None:
with _conn() as conn, conn.cursor() as cur:
cur.execute("SELECT 1 FROM pg_extension WHERE extname = 'vector'")
if not cur.fetchone():
raise RuntimeError("pgvector extension is not installed. Run: CREATE EXTENSION IF NOT EXISTS vector;")
cur.execute("SELECT to_regclass('public.rd_experiments')")
exists = bool(cur.fetchone())
if not exists:
DDL = """
CREATE TABLE IF NOT EXISTS rd_experiments (
id bigserial PRIMARY KEY NOT NULL,
experiment_name text,
rational text NOT NULL,
rational_embedding vector(384),
details text,
details_embedding vector(384),
evaluation text,
metrics jsonb,
evolved_from bigint,
start_ts timestamptz DEFAULT now() NOT NULL,
end_ts timestamptz,
git_branch text NOT NULL,
experiment_ref_id text,
session_id text,
mlruns_dir text,
status text DEFAULT 'starting' NOT NULL,
created_at timestamptz DEFAULT now() NOT NULL,
updated_at timestamptz DEFAULT now() NOT NULL
);
"""
for stmt in DDL.split(";"):
stmt = stmt.strip()
if stmt:
cur.execute(stmt)
# Idempotent backfills so pre-existing tables gain new columns.
for stmt in ["ALTER TABLE rd_experiments ADD COLUMN IF NOT EXISTS session_id text;"]:
stmt = stmt.strip()
if stmt:
cur.execute(stmt)
conn.commit()
def _search_evolved_from(text: str, limit: int = 5) -> Optional[int]:
vec = embed(text)
if not vec:
return None
lit = _vector_literal(vec)
with _conn() as conn, conn.cursor() as cur:
cur.execute(
"""
SELECT id, 1 - LEAST(
COALESCE(rational_embedding <=> %s::vector, 1),
COALESCE(details_embedding <=> %s::vector, 1)
) AS similarity
FROM rd_experiments
ORDER BY similarity DESC
LIMIT %s
""",
(lit, lit, limit),
)
rows = cur.fetchall()
for row in rows:
if row["similarity"] is not None and float(row["similarity"]) >= _MIN_SCORE:
return int(row["id"])
return None
def _search(query: str, limit: int = 10, min_score: float = _MIN_SCORE) -> List[Dict[str, Any]]:
vec = embed(query)
if not vec:
needle = f"%{query.replace('%', ' ').strip()}%"
with _conn() as conn, conn.cursor() as cur:
cur.execute(
"""
SELECT id, rational, details, git_branch, experiment_ref_id, status,
start_ts, end_ts, evaluation, session_id
FROM rd_experiments
WHERE rational ILIKE %s OR details ILIKE %s
ORDER BY id DESC LIMIT %s
""",
(needle, needle, limit),
)
return [_row_json(r) for r in cur.fetchall()]
lit = _vector_literal(vec)
with _conn() as conn, conn.cursor() as cur:
cur.execute(
"""
SELECT id, rational, details, git_branch, experiment_ref_id, status,
start_ts, end_ts, evaluation, session_id,
1 - LEAST(
COALESCE(rational_embedding <=> %s::vector, 1),
COALESCE(details_embedding <=> %s::vector, 1)
) AS similarity
FROM rd_experiments
ORDER BY similarity DESC
LIMIT %s
""",
(lit, lit, limit),
)
rows = cur.fetchall()
out = []
for r in rows:
sim = float(r.get("similarity") or 0)
if sim < min_score:
continue
r = dict(r)
r["similarity"] = sim
out.append(_row_json(r))
return out
def _start(
rational: str,
details: str = "",
evolved_from: str = "none",
experiment_name: str = "",
branch: str = "",
session_id: str = "",
) -> Dict[str, Any]:
rational = rational.strip()
details = (details or "").strip()
if not rational:
raise ValueError("--rational is required")
rational_vec = _vector_literal(embed(rational))
details_vec = _vector_literal(embed(details)) if details else None
evo: Optional[int] = None
if evolved_from == "auto":
evo = _search_evolved_from(f"{rational}\n{details}") if rational_vec or details_vec else None
elif evolved_from and evolved_from.isdigit():
evo = int(evolved_from)
with _conn() as conn, conn.cursor() as cur:
cur.execute(
"""
INSERT INTO rd_experiments
(experiment_name, rational, rational_embedding, details, details_embedding,
evolved_from, start_ts, git_branch, status, session_id)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, 'starting', %s)
RETURNING id
""",
(
experiment_name or None,
rational,
rational_vec,
details,
details_vec,
evo,
_now_iso(),
branch or "",
session_id.strip() or None,
),
)
row = cur.fetchone()
exp_id = int(row["id"])
conn.commit()
branch = branch or f"exp/{exp_id}"
with _conn() as conn, conn.cursor() as cur:
cur.execute("UPDATE rd_experiments SET git_branch = %s WHERE id = %s", (branch, exp_id))
conn.commit()
return _row_json(_get_row(exp_id) or {})
def _finish(
exp_id: int,
ref_id: str = "",
evaluation: Optional[str] = None,
metrics: Optional[str] = None,
mlruns_dir: str = "",
experiment_name: str = "",
rational: Optional[str] = None,
details: Optional[str] = None,
status: Optional[str] = None,
) -> Dict[str, Any]:
row = _get_row(exp_id)
if not row:
raise ValueError(f"experiment {exp_id} not found")
fields: List[str] = []
params: List[Any] = []
status = status or "done"
if status is None and row.get("status") in ("starting", "running"):
status = "done"
rational = (rational or row.get("rational") or "").strip()
details = (details if details is not None else row.get("details") or "").strip()
fields.append("rational = %s"); params.append(rational)
fields.append("rational_embedding = %s"); params.append(_vector_literal(embed(rational)))
fields.append("details = %s"); params.append(details)
fields.append("details_embedding = %s"); params.append(_vector_literal(embed(details)) if details else None)
if evaluation is not None:
fields.append("evaluation = %s"); params.append(evaluation.strip())
if metrics is not None:
fields.append("metrics = %s"); params.append(json.dumps(json.loads(metrics)))
if ref_id:
fields.append("experiment_ref_id = %s"); params.append(ref_id.strip())
if mlruns_dir:
fields.append("mlruns_dir = %s"); params.append(mlruns_dir.strip())
if experiment_name:
fields.append("experiment_name = %s"); params.append(experiment_name.strip())
fields.append("status = %s"); params.append(status)
fields.append("end_ts = %s"); params.append(_now_iso())
fields.append("updated_at = %s"); params.append(_now_iso())
params.append(exp_id)
with _conn() as conn, conn.cursor() as cur:
cur.execute(f"UPDATE rd_experiments SET {', '.join(fields)} WHERE id = %s", params)
conn.commit()
return _row_json(_get_row(exp_id) or {})
def _mlruns_dir(exp_name: str) -> str:
uri = (os.environ.get("DATABASE_URL") or "").strip()
if uri.startswith("postgres://"):
uri = "postgresql+psycopg://" + uri[len("postgres://") :]
if uri.startswith("postgresql://") or uri.startswith("postgresql+psycopg://"):
with _conn() as conn, conn.cursor() as cur:
cur.execute("SELECT artifact_location FROM experiments WHERE name = %s", (exp_name,))
row = cur.fetchone()
if not row:
raise RuntimeError(f"mlflow experiment {exp_name!r} not found")
return row["artifact_location"]
lake = (os.environ.get("TAC_LAKE_DIR") or "").strip()
if not lake:
raise RuntimeError("TAC_LAKE_DIR not set")
db_path = Path(lake) / "mlruns.db"
if not db_path.exists():
raise RuntimeError(f"mlruns.db not found at {db_path}")
conn = sqlite3.connect(db_path)
try:
row = conn.execute("SELECT artifact_location FROM experiments WHERE name = ?", (exp_name,)).fetchone()
finally:
conn.close()
if not row:
raise RuntimeError(f"mlflow experiment {exp_name!r} not found in {db_path}")
return row[0]
# --------------------------------------------------------------------------- git
def _parent_root() -> Path:
start = Path.cwd()
dir = start
while dir != Path(dir.anchor):
if (dir / "pnpm-workspace.yaml").exists() or (dir / "Cargo.toml").exists() or (dir / "opencode.json").exists():
return dir
dir = dir.parent
raise RuntimeError("not inside a tradeac workspace")
def _gitc(*args: str) -> subprocess.CompletedProcess:
root = _parent_root()
exp = root / "experiments"
user = (os.environ.get("GIT_USER") or "").strip()
password = (os.environ.get("GIT_PASS") or "").strip()
helper = f'!f() {{ echo "username={user}"; echo "password={password}"; }}; f'
cmd = ["git", "-C", str(exp), "-c", f"credential.helper={helper}"] + list(args)
return subprocess.run(cmd, capture_output=True, text=True)
def _git_ok(proc: subprocess.CompletedProcess) -> bool:
return proc.returncode == 0
def _git_out(proc: subprocess.CompletedProcess) -> str:
return (proc.stdout or "").strip() or (proc.stderr or "").strip()
def _require_auth() -> None:
if not (os.environ.get("GIT_REPO_URL") or "").strip() or not (os.environ.get("GIT_USER") or "").strip():
raise RuntimeError("GIT_REPO_URL / GIT_USER not set")
def _ensure_repo() -> None:
root = _parent_root()
exp = root / "experiments"
if (exp / ".git").exists():
_gitc("remote", "set-url", "origin", os.environ["GIT_REPO_URL"])
else:
_require_auth()
(root / "experiments").mkdir(parents=True, exist_ok=True)
subprocess.run(
["git", "clone", "-q", os.environ["GIT_REPO_URL"], str(exp)],
check=True, capture_output=True, text=True,
)
_gitc("config", "user.email", f"{os.environ.get('GIT_USER', '')}@tradeac.local")
_gitc("config", "user.name", os.environ.get("GIT_USER", "tradeac-agent"))
def _ensure_base(base: str = "main") -> str:
_require_auth()
_gitc("fetch", "origin", base)
if _git_ok(_gitc("rev-parse", "--verify", f"origin/{base}")):
return f"origin/{base}"
if _git_ok(_gitc("rev-parse", "--verify", base)):
return base
return base
def _fork_branch(base_ref: str, branch: str) -> str:
if _git_ok(_gitc("rev-parse", "--verify", f"origin/{branch}")):
_gitc("checkout", "-q", "-B", branch, f"origin/{branch}")
_gitc("reset", "-q", "--hard", f"origin/{branch}")
return "reused existing branch"
_gitc("fetch", "-q", "origin")
if _git_ok(_gitc("rev-parse", "--verify", f"origin/{branch}")):
_gitc("checkout", "-q", "-B", branch, f"origin/{branch}")
return "reused existing branch"
base_commit = ""
if _git_ok(_gitc("rev-parse", "--verify", f"{base_ref}^{{commit}}")):
base_commit = base_ref
elif not base_ref.startswith("origin/"):
base_commit = f"origin/{base_ref}"
if not base_commit:
raise RuntimeError(f"base '{base_ref}' not found locally or on origin")
_gitc("checkout", "-q", "-B", branch, base_commit)
return f"forked from {base_ref}"
def _snapshot_code(paths: Optional[List[str]] = None) -> str:
root = _parent_root()
exp = root / "experiments"
paths = paths or ["tac-qlib/tac_qlib/contrib", "tac-qlib/tac_qlib/data"]
parent_head = _git_out(_gitc("rev-parse", "HEAD")) or "unknown"
import shutil
shutil.rmtree(exp / "code", ignore_errors=True)
(exp / "code").mkdir(parents=True, exist_ok=True)
manifest = [f"# TradeAC custom-qlib-code snapshot (auto-generated)", f"# parent repo HEAD : {parent_head}"]
for p in paths:
manifest.append(f"# {p}")
manifest.append("# per-file hashes (git hash-object):")
for p in paths:
src = root / p
if not src.exists():
continue
dst = exp / "code" / p
dst.parent.mkdir(parents=True, exist_ok=True)
if src.is_dir():
shutil.copytree(src, dst, dirs_exist_ok=True)
for f in sorted(src.rglob("*")):
if f.is_file():
rel = str(f.relative_to(root))
h = _git_out(_gitc("hash-object", str(f)))
manifest.append(f" {h} {rel}")
else:
shutil.copy2(src, dst)
h = _git_out(_gitc("hash-object", str(src)))
manifest.append(f" {h} {p}")
(exp / "code" / "MANIFEST.txt").write_text("\n".join(manifest) + "\n")
return f"code snapshotted -> experiments/code (parent @ {parent_head[:12]})"
def _commit(message: str) -> str:
_gitc("add", "-A")
if _git_ok(_gitc("diff", "--cached", "--quiet")):
return "nothing to commit"
_gitc("commit", "-q", "-m", message)
return "committed"
def _commit_push(message: str) -> str:
result = _commit(message)
if result == "nothing to commit":
return result
branch = _git_out(_gitc("branch", "--show-current"))
_require_auth()
proc = _gitc("push", "-u", "origin", branch)
if not _git_ok(proc):
raise RuntimeError(f"push failed: {_git_out(proc)}")
return f"pushed {branch}"
def _parent_changes() -> str:
root = _parent_root()
proc = subprocess.run(["git", "-C", str(root), "status", "--porcelain"], capture_output=True, text=True)
out = (proc.stdout or "").strip()
if not out:
return "parent repo clean (no changes)"
lines = out.splitlines()
filtered = [
l for l in lines
if not l.startswith(".. experiments/")
and not l.startswith("?? experiments/")
and not l.startswith(".. tac-qlib/tac_qlib/contrib/")
and not l.startswith(".. tac-qlib/tac_qlib/data/")
and not l.startswith("?? tac-qlib/tac_qlib/contrib/")
and not l.startswith("?? tac-qlib/tac_qlib/data/")
]
expected = [l for l in lines if l.startswith(".. tac-qlib/tac_qlib/contrib/") or l.startswith(".. tac-qlib/tac_qlib/data/")]
note = ""
if expected:
note = "note: custom qlib code changed in the parent repo (contrib/data) — snapshotted to the experiment branch via trace snapshot:\n" + "\n".join(expected)
if not filtered:
return "parent repo changes limited to experiments/ and snapshotted custom qlib code (ok)" + (f"\n{note}" if note else "")
return "WARNING: unexpected parent-repo changes outside the experiments/ clone:\n" + "\n".join(filtered) + "\n→ review and revert before finishing" + (f"\n{note}" if note else "")
def slugify(text: str) -> str:
s = "".join(c for c in text.lower() if c.isalnum() or c in " -").replace(" ", "-")
return s[:40].strip("-")
def get_experiment_branch(exp_id: int) -> str:
row = _get_row(exp_id)
if not row:
raise ValueError(f"experiment {exp_id} not found")
return row["git_branch"] or f"exp/{exp_id}"
# --------------------------------------------------------------------------- MCP tools
def rd_trace_init() -> dict:
"""Ensure the traceability store + experiments git repo are ready (rd_experiments table, base main)."""
_init_db()
_ensure_repo()
base = _ensure_base("main")
return {"status": "ready", "base": base}
def rd_trace_start(
rational: str,
details: str = "",
experiment_name: str = "",
evolved_from: str = "none",
session_id: str = "",
) -> dict:
"""Open a traced experiment: insert the rd_experiments row, resolve evolved_from, fork + push the experiment branch. Pass `session_id` (the opencode chat id) so the lineage keeps a stable chat link. Returns experiment_id / branch / evolved_from / base_branch as one JSON object."""
_init_db()
_ensure_repo()
evo_id = evolved_from
if evolved_from == "auto":
evo_id = str(_search_evolved_from(rational) or "")
row = _start(rational, details, evolved_from=evo_id or "none", experiment_name=experiment_name, session_id=session_id)
exp_id = int(row["id"])
branch = f"exp/{exp_id}-{slugify(rational)}"
_gitc("checkout", "-q", "-B", branch, "main") if False else None
# fork from the predecessor branch (or main)
base_branch = "main"
if evo_id and evo_id.isdigit():
base_branch = get_experiment_branch(int(evo_id))
base_ref = _ensure_base(base_branch)
_fork_branch(base_ref, branch)
_snapshot_code()
_gitc("checkout", "-q", "-B", branch, branch) if False else None
# persist branch on the row
with _conn() as conn, conn.cursor() as cur:
cur.execute("UPDATE rd_experiments SET git_branch = %s WHERE id = %s", (branch, exp_id))
conn.commit()
_commit_push(f"start experiment {exp_id} ({branch})")
return {"experiment_id": exp_id, "branch": branch, "evolved_from": evo_id or "none", "base_branch": base_branch}
def rd_trace_finish(
experiment_id: int,
ref_id: str = "",
evaluation: str = "",
metrics: str = "",
mlruns_dir: str = "",
experiment_name: str = "",
) -> dict:
"""Close a traced experiment: update the row (link the mlflow run, metrics/evaluation/end_ts), snapshot code, commit + push the branch. Returns the updated row."""
_finish(experiment_id, ref_id=ref_id, evaluation=evaluation or None, metrics=metrics or None, mlruns_dir=mlruns_dir, experiment_name=experiment_name)
branch = get_experiment_branch(experiment_id)
_gitc("checkout", "-q", "-B", branch, branch)
_snapshot_code()
_commit_push(f"finish experiment {experiment_id} ({branch})")
return {"experiment_id": experiment_id, "branch": branch, "row": _row_json(_get_row(experiment_id) or {})}
def rd_trace_commit(experiment_id: int, message: str = "wip") -> dict:
"""Commit the current experiment branch state (no push)."""
branch = get_experiment_branch(experiment_id)
_gitc("checkout", "-q", "-B", branch, branch)
result = _commit(f"exp {experiment_id}: {message}")
return {"experiment_id": experiment_id, "branch": branch, "result": result}
def rd_trace_snapshot(experiment_id: int, paths: str = "") -> dict:
"""Snapshot custom qlib contrib/data code onto the experiment branch (default contrib+data)."""
branch = get_experiment_branch(experiment_id)
_gitc("checkout", "-q", "-B", branch, branch)
path_list = [p.strip() for p in paths.split(",") if p.strip()] if paths else None
msg = _snapshot_code(path_list)
_commit_push(f"exp {experiment_id}: snapshot custom qlib code")
return {"experiment_id": experiment_id, "branch": branch, "result": msg}
def rd_trace_guard() -> dict:
"""Check the parent repo for unexpected changes outside the experiments clone."""
return {"parent_changes": _parent_changes()}
def rd_trace_search(query: str, limit: int = 10, min_score: float = _MIN_SCORE) -> dict:
"""Semantic search over experiment rationals/details (pgvector, falls back to ILIKE)."""
return {"results": _search(query, limit=limit, min_score=min_score)}
def rd_trace_get(experiment_id: int) -> dict:
"""Return one traced experiment row."""
row = _get_row(experiment_id)
if not row:
raise ValueError(f"experiment {experiment_id} not found")
return _row_json(row)
def rd_trace_list(limit: int = 20) -> dict:
"""List traced experiments (newest first)."""
return {"experiments": [_row_json(r) for r in _all_rows(limit)]}
def rd_trace_mlruns_dir(experiment_name: str) -> dict:
"""Resolve the mlflow artifact location for an experiment name."""
return {"mlruns_dir": _mlruns_dir(experiment_name)}
def register_trace_tools(server) -> None:
"""Attach all rd_trace_* tools to an MCPServer instance (called by rd_server.main())."""
for fn in (
rd_trace_init,
rd_trace_start,
rd_trace_finish,
rd_trace_commit,
rd_trace_snapshot,
rd_trace_guard,
rd_trace_search,
rd_trace_get,
rd_trace_list,
rd_trace_mlruns_dir,
):
server.tool(structured_output=False)(fn)
+59
View File
@@ -0,0 +1,59 @@
"""Embed text via the self-hosted infinity embedding API (used by tac_qlib.trace).
Mirrors the standalone `embed.py` in the tac-qlib-custom skill lib so the trace
MCP tools can embed rational/details without shelling out.
"""
from __future__ import annotations
import base64
import json
import os
import urllib.request
EMBEDDING_MODEL = "michaelfeil/bge-small-en-v1.5"
MAX_TOKENS = 512
CHARS_PER_TOKEN = 4
def estimate_tokens(text: str) -> int:
return max(1, -(-len(text) // CHARS_PER_TOKEN))
def embed(text: str, timeout: int = 40) -> list[float] | None:
base_url = (os.environ.get("EMBEDDING_API_BASE_URL") or "").strip()
api_key = (os.environ.get("EMBEDDING_API_KEY") or "").strip()
if not base_url or not api_key:
return None
if estimate_tokens(text) > MAX_TOKENS:
raise ValueError(
f"text is ~{estimate_tokens(text)} tokens, exceeding the {MAX_TOKENS}-token embedding "
"context. Write a <=512-token summary of the experiment and embed that instead."
)
body = json.dumps({"model": EMBEDDING_MODEL, "input": text}).encode("utf-8")
req = urllib.request.Request(
base_url,
data=body,
headers={
"accept": "application/json",
"Content-Type": "application/json",
},
)
user, _, password = api_key.partition(":")
cred = base64.b64encode(f"{user}:{password}".encode("utf-8")).decode("ascii")
req.add_header("Authorization", f"Basic {cred}")
with urllib.request.urlopen(req, timeout=timeout) as resp:
payload = json.loads(resp.read().decode("utf-8"))
data = payload.get("data") if isinstance(payload, dict) else None
if isinstance(data, list) and data and isinstance(data[0], dict):
emb = data[0].get("embedding")
if isinstance(emb, list) and emb:
return [float(v) for v in emb]
embeddings = payload.get("embeddings") if isinstance(payload, dict) else None
if isinstance(embeddings, list) and embeddings and isinstance(embeddings[0], list):
return [float(v) for v in embeddings[0]]
raise RuntimeError(f"unexpected embedding response shape: {str(payload)[:300]}")