Files
worldquant-alpha-system/backend/app/research/evaluations.py
T

174 lines
7.1 KiB
Python
Raw Normal View History

"""Snapshot-based research assessment. Model advice cannot replace deterministic findings."""
import math
from fastapi import HTTPException
from sqlalchemy import select
from ..backtests.service import Backtests, uid
from ..models import Alpha, ResearchEvaluation, ResearchExperiment, SelfCorrelation
from ..platform_checks import split_checks, submission_limits
from .experiments import Experiments
from .serialization import encode_snapshot as jsonable_encoder
def assess(snapshot, rules):
metrics = snapshot.get("is") or {}
evidence, missing, failed = [], [], []
for key, bound, direction in (
("sharpe", rules.sharpe_min, "min"),
("fitness", rules.fitness_min, "min"),
("turnover", rules.turnover_max, "max"),
):
value = metrics.get(key)
if type(value) not in (int, float) or not math.isfinite(value):
missing.append(key)
value, status = None, "unknown"
else:
status = "pass" if (value >= bound if direction == "min" else value <= bound) else "block"
if status == "block":
failed.append(key)
evidence.append(
{"metric": key, "value": value, "bound": bound, "direction": direction, "status": status}
)
raw_checks = metrics.get("checks") or []
checks, _ = split_checks(raw_checks)
if not checks:
missing.append("platform_checks")
for check in checks:
if isinstance(check, dict) and check.get("result") == "FAIL":
failed.append(f"platform:{check.get('name', 'unknown')}")
unknown_checks = [check for check in checks if not isinstance(check, dict) or check.get("result") not in ("PASS", "FAIL")]
if unknown_checks:
missing.append("unresolved_platform_checks")
return {
"verdict": "block" if failed else "review" if missing else "pass",
"evidence": evidence,
"failed": failed,
"missing": missing,
"existing_platform_checks": raw_checks,
"submission_limits": submission_limits(raw_checks),
"meaning": "本地研究筛选结果,不是官方提交资格",
}
class Evaluations:
def __init__(self, db):
self.db = db
async def create(self, body):
records = []
if body.alpha_id:
alpha = await self.db.get(Alpha, body.alpha_id)
if not alpha:
raise HTTPException(404, "Alpha 尚未同步")
records.append(
{
"alpha_id": alpha.id,
"snapshot": alpha.raw,
"observed_at": alpha.synced_at,
"client_item_id": None,
}
)
else:
run = await Backtests(self.db).run(body.backtest_run_id)
if body.experiment_id and run["source"].get("research_id") != body.experiment_id:
raise HTTPException(422, "回测运行不属于指定研究实验")
for offset in range(0, run["total"], 100):
page = await Backtests(self.db).results(body.backtest_run_id, 100, offset)
for item in page["items"]:
records.append(
{
"alpha_id": item["alpha_id"],
"snapshot": item["result"]["snapshot"] if item["result"] else {},
"observed_at": item["result"]["observed_at"] if item["result"] else None,
"client_item_id": item["client_item_id"],
"item_id": item["id"],
"error": item["error"],
"complete": bool(item["result"] and item["result"]["complete"]),
"expression": item["expression"],
"settings": item["settings"],
}
)
if body.experiment_id and not await self.db.get(ResearchExperiment, body.experiment_id):
raise HTTPException(404, "研究实验不存在")
findings = []
for record in records:
correlation = (
await self.db.get(SelfCorrelation, record["alpha_id"]) if record["alpha_id"] else None
)
finding = assess(record["snapshot"], body.rules)
if record.get("error") or record.get("complete") is False:
finding["missing"].append("backtest_error")
if finding["verdict"] == "pass":
finding["verdict"] = "review"
findings.append(
{
**record,
**finding,
"local_correlation": {
"result": correlation.result,
"stale": correlation.stale,
"calculated_at": correlation.calculated_at,
}
if correlation
else None,
}
)
report = jsonable_encoder(
{
"rules": body.rules.model_dump(),
"records": findings,
"backtest_run_id": body.backtest_run_id,
"experiment": await Experiments(self.db).get(body.experiment_id)
if body.experiment_id
else None,
"model_advice": None,
"verdict": "block"
if any(r["verdict"] == "block" for r in findings)
else "review"
if not findings or any(r["verdict"] == "review" for r in findings)
else "pass",
}
)
row = ResearchEvaluation(
id=uid(), alpha_id=body.alpha_id, experiment_id=body.experiment_id, report=report
)
self.db.add(row)
await self.db.flush()
return await self.get(row.id)
async def get(self, evaluation_id):
row = await self.db.get(ResearchEvaluation, evaluation_id)
if not row:
raise HTTPException(404, "评估报告不存在")
return jsonable_encoder(
{key: getattr(row, key) for key in ("id", "alpha_id", "experiment_id", "report", "created_at")}
)
async def list(self, alpha_id=None, experiment_id=None, limit=25, offset=0):
query = select(ResearchEvaluation)
if alpha_id:
query = query.where(ResearchEvaluation.alpha_id == alpha_id)
if experiment_id:
query = query.where(ResearchEvaluation.experiment_id == experiment_id)
rows = await self.db.scalars(
query.order_by(ResearchEvaluation.created_at.desc()).limit(limit).offset(offset)
)
return {"items": [await self.get(row.id) for row in rows]}
async def add_advice(self, evaluation_id, advice, model_evidence):
original = await self.get(evaluation_id)
report = {
**original["report"],
"model_advice": advice,
"model_evidence": model_evidence,
"previous_evaluation_id": evaluation_id,
}
row = ResearchEvaluation(
id=uid(), alpha_id=original["alpha_id"], experiment_id=original["experiment_id"], report=report
)
self.db.add(row)
await self.db.flush()
return await self.get(row.id)