Files
myquant/finance/agents/selection_agent.py
T
Simon 73d191b43a feat: 量化引擎加固 — 新增测试 + 数据/因子/回测层优化
- 新增 finance/tests/ 6 个测试套件(agents/backtest/dao_upsert/factors/features/fundamental_lookahead)
- 数据层: data_manager / dao 优化,新增 upsert 逻辑
- 因子层: 基本面因子抽象定位 _mapping、ROE/PE/PB 重构
- 回测层: vectorbt/engine 大改动(251 行),report 增强
- ML 层: features/backtest_integration 特征工程与回测优化
- CLI: agent_cli 重构
- config/settings 扩充配置项
2026-08-31 14:01:06 +08:00

212 lines
7.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
SelectionAgent — 多因子股票打分。
综合技术因子 + 基本面因子 + ML 预测值 → 股票综合得分排序。
"""
import numpy as np
import pandas as pd
from agents.base import BaseAgent
from config.settings import (
SELECTION_CORE_FACTORS, SELECTION_UNIVERSE_SIZE,
SELECTION_SCORE_LIMIT, SELECTION_WINSORIZE_ZSCORE,
)
class SelectionAgent(BaseAgent):
"""股票打分 Agent。"""
name = "Selection"
description = "多因子股票打分与推荐"
def execute(
self,
date: str | None = None,
ts_codes: list[str] | None = None,
top_n: int = 15,
weighting: str = "equal",
) -> dict:
"""
对股票池打分排序。
参数:
date: 目标日期,None=最新
ts_codes: 股票池,None=指数成分股
top_n: 返回 Top N
weighting: 'equal' | 'ic_weighted' | 'ml'
返回:
{"date": ..., "top_picks": [...], "score_df": DataFrame}
"""
from factors.registry import get_factor, list_factors
# 确定股票池
if ts_codes is None:
if self.sent:
ts_codes = self.sent.get_scope_stocks()
else:
stocks = self.dm.get_stock_list()
ts_codes = list(stocks.index[:SELECTION_UNIVERSE_SIZE])
if not ts_codes:
return {"date": date or self._today(), "top_picks": [], "score_df": pd.DataFrame()}
# 选择核心因子(覆盖多个维度,减少计算量;参数来自配置中心)
core_factors = list(SELECTION_CORE_FACTORS)
factor_objects = [get_factor(n) for n in core_factors]
self.log("打分 {} 只股票 (权重={})".format(len(ts_codes), weighting))
# 1. 筛选有 DB 缓存的股票(Orchestrator 已在 Step1 同步了全量范围)
from database.dao import get_latest_trade_date
available = []
for ts_code in sorted(ts_codes):
if get_latest_trade_date(ts_code):
available.append(ts_code)
self.log("缓存命中: {}/{} ({:.1f}%)".format(
len(available), len(ts_codes),
len(available) / len(ts_codes) * 100 if ts_codes else 0))
# 2. 打分(限制上限防止单次太慢)
score_limit = min(len(available), SELECTION_SCORE_LIMIT)
scores = {}
valid_count = 0
for i, ts_code in enumerate(available[:score_limit]):
try:
score = self._score_stock(ts_code, factor_objects, date, weighting)
if score is not None:
scores[ts_code] = score
valid_count += 1
except Exception:
continue
if (i + 1) % 50 == 0:
self.log(" 进度: {}/{}".format(i + 1, score_limit))
self.log("有效评分: {}/{}".format(len(scores), score_limit))
if not scores:
return {"date": date or self._today(), "top_picks": [], "score_df": pd.DataFrame()}
# 排序
sorted_scores = sorted(scores.items(), key=lambda x: x[1], reverse=True)
top = sorted_scores[:top_n]
# 取股票名称(ts_code 是 index)
stock_list = self.dm.get_stock_list() if self.dm else pd.DataFrame()
name_map = {}
if not stock_list.empty and "name" in stock_list.columns:
name_map = dict(zip(stock_list.index, stock_list["name"]))
top_picks = []
for ts_code, score in top:
# 尝试多种格式匹配名称
code_no_suffix = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
name = name_map.get(ts_code, name_map.get(code_no_suffix, ts_code))
top_picks.append({
"ts_code": ts_code,
"name": name,
"score": round(score, 4),
})
score_df = pd.DataFrame(
{"ts_code": list(scores.keys()), "score": list(scores.values())}
).sort_values("score", ascending=False).reset_index(drop=True)
return {
"date": date or self._today(),
"top_picks": top_picks,
"score_df": score_df,
"universe_size": len(ts_codes),
"valid_scores": valid_count,
}
def _filter_cached_stocks(self, ts_codes: list[str], limit: int = 100) -> list[str]:
"""筛选有 DB 日线缓存的股票,避免逐个调用 AkShare。"""
from database.dao import get_latest_trade_date
cached = []
for ts_code in ts_codes[:limit]:
latest = get_latest_trade_date(ts_code)
if latest:
cached.append(ts_code)
return cached
def _score_stock(
self,
ts_code: str,
factors: list,
date: str | None,
weighting: str,
) -> float | None:
"""对单只股票打分。"""
daily = self.dm.get_daily(ts_code)
if daily is None or daily.empty:
return None
daily = daily.set_index("trade_date").sort_index()
factor_df = self.fe.compute(ts_code, factors)
if factor_df is None or factor_df.empty:
return None
if date and date in factor_df.index:
row = factor_df.loc[date]
else:
row = factor_df.iloc[-1] # 最新一天
if row.isna().all():
return None
if weighting == "ml" and self.ml_models:
return self._score_ml(ts_code, factor_df, date)
# 等权打分:标准化因子值后求和
row_clean = row.dropna()
if len(row_clean) < 3:
return None
# z-score 标准化(用历史均值/标准差),可选的去极值避免单股离群主导 Top
cols = row_clean.index
mu = factor_df[cols].mean()
std = factor_df[cols].std().replace(0, 1)
z = (row_clean - mu) / std
if SELECTION_WINSORIZE_ZSCORE:
z = z.clip(-3, 3)
return float(z.mean())
def _score_ml(self, ts_code: str, factor_df: pd.DataFrame, date: str | None) -> float | None:
"""ML 模型打分。
要求注入名为 feature_engine 的、已用训练集 fit 过的 FeatureEngine,
以及 ml_models(已训练模型)。两者缺一时明确退出,而不是沿用旧的
未 fit 引擎静默失败。
"""
fe = self.feature_engine
if fe is None:
self.log("ML 打分需要注入 feature_engine(已 fit),当前未提供,跳过 ML 打分")
return None
if not self.ml_models:
self.log("ML 打分需要 ml_models(已训练),当前为空,跳过 ML 打分")
return None
daily = self.dm.get_daily(ts_code) if self.dm else None
if daily is None or daily.empty:
return None
daily = daily.set_index("trade_date")
try:
X, _ = fe.build(factor_df, daily, fit=False)
if X is None or X.empty:
return None
if date and date in X.index:
X = X.loc[[date]]
else:
X = X.iloc[[-1]]
model = self.ml_models.get("lightgbm") or list(self.ml_models.values())[0]
pred = model.predict(X)
return float(pred.iloc[0]) if len(pred) > 0 else None
except Exception as e:
# 不再静默返回 None:记录原因,便于定位预测路径问题
self.log("[WARN] ML 打分失败 ({}): {}".format(ts_code, e))
return None