118 lines
4.4 KiB
Python
118 lines
4.4 KiB
Python
"""因子批量评估 API:runs/leaderboard/detail/submit."""
|
|
from fastapi import APIRouter, HTTPException, Depends, Query
|
|
|
|
from .schemas import FactorBatchEvalRequest
|
|
from .routes import verify_token
|
|
|
|
router = APIRouter(dependencies=[Depends(verify_token)])
|
|
|
|
_eval_db_path: str | None = None
|
|
|
|
_SORT_KEYS = {
|
|
"abs_icir": lambda p: abs(p.get("icir") or 0.0),
|
|
"ic_mean": lambda p: p.get("ic_mean") if p.get("ic_mean") is not None else float("-inf"),
|
|
"t_stat": lambda p: abs(p.get("t_stat") or 0.0),
|
|
"ls_annual": lambda p: p.get("ls_annual") if p.get("ls_annual") is not None else float("-inf"),
|
|
}
|
|
|
|
|
|
def set_eval_db_path(path: str) -> None:
|
|
global _eval_db_path
|
|
_eval_db_path = path
|
|
|
|
|
|
def _db() -> str:
|
|
from sanguo_factor.eval_store import default_eval_db_path
|
|
return _eval_db_path or default_eval_db_path()
|
|
|
|
|
|
def _latest_run_id(db: str) -> str | None:
|
|
from sanguo_factor.eval_store import list_runs
|
|
runs = list_runs(db)
|
|
return runs[0]["run_id"] if runs else None
|
|
|
|
|
|
@router.get("/factor/eval/runs")
|
|
def eval_runs():
|
|
from sanguo_factor.eval_store import list_runs
|
|
return {"runs": list_runs(_db())}
|
|
|
|
|
|
@router.get("/factor/eval/leaderboard")
|
|
def eval_leaderboard(
|
|
run_id: str | None = None,
|
|
period: str = Query("1", pattern="^(1|5|10)$"),
|
|
category: str | None = None,
|
|
search: str | None = None,
|
|
sort: str = Query("abs_icir", pattern="^(abs_icir|ic_mean|t_stat|ls_annual)$"),
|
|
order: str = Query("desc", pattern="^(asc|desc)$"),
|
|
):
|
|
from sanguo_factor.eval_store import get_rows
|
|
db = _db()
|
|
rid = run_id or _latest_run_id(db)
|
|
if not rid:
|
|
return {"tiles": {"factors_total": 0, "effective": 0, "watch": 0,
|
|
"top_ls_annual": None, "top_ls_factor": None}, "rows": []}
|
|
rows = get_rows(db, rid, category=category, search=search)
|
|
|
|
flat: list[dict] = []
|
|
for r in rows:
|
|
m = r.get("metrics") or {}
|
|
p = m.get(period) or {}
|
|
if "error" in m and not p:
|
|
flat.append({"factor": r["factor"], "category": r["category"],
|
|
"expression": r["expression"], "error": m["error"],
|
|
"ic_mean": None, "icir": None, "t_stat": None, "win_rate": None,
|
|
"ls_annual": None, "turnover": m.get("turnover"),
|
|
"conclusion": "eliminated", "monthly_ic": []})
|
|
continue
|
|
flat.append({
|
|
"factor": r["factor"], "category": r["category"], "expression": r["expression"],
|
|
"ic_mean": p.get("ic_mean"), "icir": p.get("icir"), "t_stat": p.get("t_stat"),
|
|
"win_rate": p.get("win_rate"), "ls_annual": p.get("ls_annual"),
|
|
"turnover": m.get("turnover"), "conclusion": p.get("conclusion", "eliminated"),
|
|
"monthly_ic": p.get("monthly_ic", []),
|
|
})
|
|
|
|
keyfn = _SORT_KEYS[sort]
|
|
flat.sort(key=lambda x: (keyfn(x) is not None, keyfn(x)), reverse=(order == "desc"))
|
|
for i, row in enumerate(flat, 1):
|
|
row["rank"] = i
|
|
|
|
scored = [x for x in flat if x.get("ic_mean") is not None]
|
|
top_ls = max(scored, key=lambda x: (x.get("ls_annual") or float("-inf")), default=None)
|
|
tiles = {
|
|
"factors_total": len(flat),
|
|
"effective": sum(1 for x in flat if x["conclusion"] == "effective"),
|
|
"watch": sum(1 for x in flat if x["conclusion"] == "watch"),
|
|
"top_ls_annual": top_ls["ls_annual"] if top_ls else None,
|
|
"top_ls_factor": top_ls["factor"] if top_ls else None,
|
|
}
|
|
return {"tiles": tiles, "rows": flat}
|
|
|
|
|
|
@router.get("/factor/eval/detail")
|
|
def eval_detail(run_id: str | None = None, factor: str = ""):
|
|
from sanguo_factor.eval_store import get_detail
|
|
db = _db()
|
|
rid = run_id or _latest_run_id(db)
|
|
d = get_detail(db, rid, factor) if rid else None
|
|
if d is None:
|
|
raise HTTPException(status_code=404, detail="factor not found in run")
|
|
return d
|
|
|
|
|
|
@router.post("/factor/eval/submit")
|
|
async def eval_submit(req: FactorBatchEvalRequest):
|
|
if not req.categories and not req.factors:
|
|
raise HTTPException(status_code=400, detail="categories 与 factors 至少给一个")
|
|
from .routes import get_orchestrator
|
|
orch = get_orchestrator()
|
|
if orch is None:
|
|
raise HTTPException(status_code=503, detail="orchestrator 未就绪")
|
|
task_id = await orch.submit_batch_eval(
|
|
factor_names=req.factors, categories=req.categories,
|
|
start=req.start, end=req.end, symbols=req.symbols or None, label=req.label,
|
|
)
|
|
return {"task_id": task_id}
|