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 扩充配置项
This commit is contained in:
@@ -25,6 +25,7 @@ class BaseAgent(ABC):
|
||||
opt: OptunaEngine
|
||||
sent: SentimentEngine
|
||||
ml_models: dict[str, BaseModel]
|
||||
feature_engine: FeatureEngine(已用模型训练集 fit 过,ML 打分需要)
|
||||
"""
|
||||
self.dm = engines.get("dm")
|
||||
self.fe = engines.get("fe")
|
||||
@@ -32,6 +33,7 @@ class BaseAgent(ABC):
|
||||
self.opt = engines.get("opt")
|
||||
self.sent = engines.get("sent")
|
||||
self.ml_models = engines.get("ml_models", {})
|
||||
self.feature_engine = engines.get("feature_engine")
|
||||
|
||||
@abstractmethod
|
||||
def execute(self, **kwargs) -> dict:
|
||||
|
||||
+117
-72
@@ -4,6 +4,7 @@ AgentOrchestrator — Agent 编排器。
|
||||
管理所有 Agent 的生命周期、执行顺序、结果传递。
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
|
||||
from agents.research_agent import ResearchAgent
|
||||
@@ -11,6 +12,10 @@ from agents.selection_agent import SelectionAgent
|
||||
from agents.risk_agent import RiskAgent
|
||||
from agents.report_agent import ReportAgent
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TOTAL_STEPS = 5
|
||||
|
||||
|
||||
class AgentOrchestrator:
|
||||
"""
|
||||
@@ -27,112 +32,153 @@ class AgentOrchestrator:
|
||||
self.agents: dict = {}
|
||||
self._last_results: dict = {}
|
||||
|
||||
def _log_step(self, current: int, title: str) -> None:
|
||||
"""统一步骤日志:分母固定为 TOTAL_STEPS。"""
|
||||
print("\n[Step {}/{}] {}...".format(current, TOTAL_STEPS, title))
|
||||
|
||||
def _run_step(self, step_no, key, title, fn, results, on_ok=None):
|
||||
"""执行单个步骤并隔离异常。
|
||||
|
||||
key: 结果字典的键(英文,向前兼容如 risk/selection/report),
|
||||
title: 步骤展示名(中文,仅用于日志)。
|
||||
任一子步骤抛错都会被捕获:记录 results[key]={"error": ...},
|
||||
让失败不中断整体流程;on_ok 用于成功后的副作用(打印等)。
|
||||
"""
|
||||
self._log_step(step_no, title)
|
||||
try:
|
||||
value = fn()
|
||||
results[key] = value
|
||||
if on_ok:
|
||||
on_ok(value)
|
||||
return value
|
||||
except Exception as e:
|
||||
logger.exception("[Orchestrator] 步骤 %s 失败: %s", title, e)
|
||||
results[key] = {"error": str(e)}
|
||||
print(" [SKIP] 步骤 {} 失败,继续后续步骤: {}".format(title, e))
|
||||
return None
|
||||
|
||||
def setup(self):
|
||||
"""注册所有 Agent。"""
|
||||
self.agents["research"] = ResearchAgent(**self.engines)
|
||||
self.agents["selection"] = SelectionAgent(**self.engines)
|
||||
self.agents["risk"] = RiskAgent(**self.engines)
|
||||
self.agents["report"] = ReportAgent(**self.engines)
|
||||
print(f"[Orchestrator] 已注册 {len(self.agents)} 个 Agent: {list(self.agents)}")
|
||||
logger.info("[Orchestrator] 已注册 %s 个 Agent: %s",
|
||||
len(self.agents), list(self.agents))
|
||||
|
||||
# ── 每日流程 ──────────────────────────────────────────
|
||||
|
||||
def run_daily(self, date: str | None = None) -> dict:
|
||||
"""
|
||||
每日任务流:
|
||||
每日任务流(5 步,任一步失败不中断整体):
|
||||
|
||||
1. 同步数据
|
||||
1. 同步数据(增量同步已缓存股票)
|
||||
2. 风险评估
|
||||
3. 股票打分
|
||||
4. 生成日报
|
||||
4. 情绪因子
|
||||
5. 生成日报
|
||||
"""
|
||||
date = date or datetime.now().strftime("%Y%m%d")
|
||||
print(f"\n{'='*60}")
|
||||
print(f"[Orchestrator] 每日流程 — {date}")
|
||||
logger.info("[Orchestrator] 每日流程 — %s", date)
|
||||
print(f"{'='*60}")
|
||||
|
||||
results = {"date": date}
|
||||
results: dict = {"date": date}
|
||||
|
||||
# Step 1: 增量同步已缓存股票的最新行情
|
||||
print("\n[Step 1/4] 同步行情...")
|
||||
dm = self.engines.get("dm")
|
||||
sent = self.engines.get("sent")
|
||||
from database.dao import get_latest_trade_date
|
||||
|
||||
scope_stocks = []
|
||||
if sent:
|
||||
scope_stocks = sent.get_scope_stocks()
|
||||
if not scope_stocks and dm:
|
||||
scope_stocks = list(dm.get_stock_list().index[:100])
|
||||
print(" 范围: {} 只股票".format(len(scope_stocks)))
|
||||
def _step_sync():
|
||||
scope_stocks = []
|
||||
if sent:
|
||||
scope_stocks = sent.get_scope_stocks()
|
||||
if not scope_stocks and dm:
|
||||
scope_stocks = list(dm.get_stock_list().index[:100])
|
||||
print(" 范围: {} 只股票".format(len(scope_stocks)))
|
||||
|
||||
# 分类:已缓存(增量更新),未缓存(统计跳过)
|
||||
cached = [c for c in scope_stocks if get_latest_trade_date(c)]
|
||||
uncached = len(scope_stocks) - len(cached)
|
||||
print(" 已缓存: {} 只 (增量更新), 未缓存: {} 只 (跳过,需首次批量预热)".format(len(cached), uncached))
|
||||
cached = [c for c in scope_stocks if get_latest_trade_date(c)]
|
||||
uncached = len(scope_stocks) - len(cached)
|
||||
print(" 已缓存: {} 只 (增量更新), 未缓存: {} 只 (需首次预热)".format(len(cached), uncached))
|
||||
|
||||
synced = 0
|
||||
for i, ts_code in enumerate(cached):
|
||||
try:
|
||||
n = dm.sync_daily(ts_code)
|
||||
synced += n
|
||||
except Exception:
|
||||
continue
|
||||
if (i + 1) % 100 == 0:
|
||||
print(" [同步] 进度: {}/{}".format(i + 1, len(cached)))
|
||||
print(" 同步完成: {} 条新数据 (已缓存{}/全量{})".format(synced, len(cached), len(scope_stocks)))
|
||||
synced = 0
|
||||
for i, ts_code in enumerate(cached):
|
||||
try:
|
||||
n = dm.sync_daily(ts_code)
|
||||
synced += n
|
||||
except Exception as e:
|
||||
logger.debug("同步 %s 失败: %s", ts_code, e)
|
||||
continue
|
||||
if (i + 1) % 100 == 0:
|
||||
logger.info("[同步] 进度: %s/%s", i + 1, len(cached))
|
||||
if uncached > 0:
|
||||
print(" 提示: {} 只股票未缓存,运行 'agent_cli.py warmup' 首次批量预热".format(uncached))
|
||||
return {"synced": synced, "cached": len(cached), "uncached": uncached,
|
||||
"universe_size": len(scope_stocks), "scope_stocks": scope_stocks}
|
||||
|
||||
if uncached > 0:
|
||||
print(" 提示: {} 只股票未缓存,运行 'agent_cli.py warmup' 首次批量预热".format(uncached))
|
||||
sync_info = self._run_step(1, "sync", "同步行情", _step_sync, results, on_ok=lambda r: print(
|
||||
" 同步完成: {} 条新数据 (已缓存{}/全量{})".format(
|
||||
r["synced"], r["cached"], r["universe_size"])))
|
||||
|
||||
# Step 2: 风险评估
|
||||
print("\n[Step 2/4] 风险评估...")
|
||||
risk = self.agents["risk"].execute()
|
||||
results["risk"] = risk
|
||||
print(f" 风险: {risk['risk_level']}, 仓位: {risk['target_exposure']:.0%}")
|
||||
def _risk():
|
||||
risk = self.agents["risk"].execute()
|
||||
risk_level = risk.get("risk_level", "?")
|
||||
exposure = risk.get("target_exposure", 0)
|
||||
print(" 风险: {}, 仓位: {:.0%}".format(risk_level, exposure))
|
||||
return risk
|
||||
risk = self._run_step(2, "risk", "风险评估", _risk, results)
|
||||
|
||||
# Step 3: 选股打分
|
||||
print("\n[Step 3/4] 股票打分...")
|
||||
selection = self.agents["selection"].execute(date=date, top_n=15)
|
||||
results["selection"] = selection
|
||||
top = selection.get("top_picks", [])
|
||||
if top:
|
||||
print(" Top 5: {}".format(", ".join(p['ts_code'] for p in top[:5])))
|
||||
def _selection():
|
||||
selection = self.agents["selection"].execute(date=date, top_n=15)
|
||||
top = selection.get("top_picks", [])
|
||||
if top:
|
||||
print(" Top 5: {}".format(", ".join(p['ts_code'] for p in top[:5])))
|
||||
return selection
|
||||
selection = self._run_step(3, "selection", "股票打分", _selection, results)
|
||||
|
||||
# Step 4: 情绪因子
|
||||
print("\n[Step 4/5] 情绪因子...")
|
||||
sentiment_df = None
|
||||
sent_eng = self.engines.get("sent")
|
||||
if sent_eng:
|
||||
try:
|
||||
# 用范围内第一只有缓存的股票计算情绪因子
|
||||
ref_code = cached[0] if cached else "000001.SZ"
|
||||
sentiment_df = sent_eng.compute(ref_code, max_news=30)
|
||||
if sentiment_df is not None and not sentiment_df.empty:
|
||||
valid = sentiment_df.dropna(how="all")
|
||||
print(" {}: {} 个因子, {} 个有效交易日".format(ref_code, sentiment_df.shape[1], len(valid)))
|
||||
else:
|
||||
print(" (无有效情绪数据)")
|
||||
except Exception as e:
|
||||
print(" [SKIP] 情绪因子计算失败: {}".format(e))
|
||||
else:
|
||||
print(" (SentimentEngine 未配置)")
|
||||
results["sentiment"] = sentiment_df
|
||||
def _sentiment():
|
||||
sentiment_df = None
|
||||
sent_eng = self.engines.get("sent")
|
||||
if not sent_eng:
|
||||
print(" (SentimentEngine 未配置)")
|
||||
return None
|
||||
# 用范围内第一只有缓存的股票计算情绪因子;无缓存信息时回退沪市指数代码
|
||||
scope = (sync_info or {}).get("scope_stocks") or []
|
||||
ref_code = scope[0] if scope else "000001.SZ"
|
||||
sentiment_df = sent_eng.compute(ref_code, max_news=30)
|
||||
if sentiment_df is not None and not sentiment_df.empty:
|
||||
valid = sentiment_df.dropna(how="all")
|
||||
print(" {}: {} 个因子, {} 个有效交易日".format(
|
||||
ref_code, sentiment_df.shape[1], len(valid)))
|
||||
else:
|
||||
print(" (无有效情绪数据)")
|
||||
return sentiment_df
|
||||
sentiment_df = self._run_step(4, "sentiment", "情绪因子", _sentiment, results)
|
||||
|
||||
# Step 5: 生成日报
|
||||
print("\n[Step 5/5] 生成日报...")
|
||||
report = self.agents["report"].execute(
|
||||
date=date,
|
||||
selection_result=selection,
|
||||
risk_result=risk,
|
||||
sentiment_result=sentiment_df,
|
||||
)
|
||||
results["report"] = report
|
||||
print(f" 日报: {report['report_path']}")
|
||||
def _report():
|
||||
rep = self.agents["report"].execute(
|
||||
date=date,
|
||||
selection_result=selection,
|
||||
risk_result=risk,
|
||||
sentiment_result=sentiment_df,
|
||||
)
|
||||
print(" 日报: {}".format(rep.get("report_path", "?")))
|
||||
return rep
|
||||
report = self._run_step(5, "report", "生成日报", _report, results)
|
||||
|
||||
self._last_results = results
|
||||
print(f"\n{'='*60}")
|
||||
print(f"[Orchestrator] 每日流程完成")
|
||||
print("[Orchestrator] 每日流程完成")
|
||||
failed = [k for k, v in results.items() if isinstance(v, dict) and "error" in v]
|
||||
if failed:
|
||||
print("[Orchestrator] 有步骤未完成: {}".format(", ".join(failed)))
|
||||
else:
|
||||
print("[Orchestrator] 所有步骤均成功")
|
||||
print(f"{'='*60}\n")
|
||||
return results
|
||||
|
||||
@@ -140,16 +186,15 @@ class AgentOrchestrator:
|
||||
|
||||
def run_research_cycle(self, ts_codes: list[str] | None = None) -> dict:
|
||||
"""
|
||||
研究周期:
|
||||
研究周期(1 步):
|
||||
|
||||
1. 因子发现与评估
|
||||
2. 更新 IC 权重
|
||||
1. 因子发现与评估(IC / IC_IR / 分层收益)
|
||||
"""
|
||||
print(f"\n{'='*60}")
|
||||
print(f"[Orchestrator] 研究周期")
|
||||
logger.info("[Orchestrator] 研究周期")
|
||||
print(f"{'='*60}")
|
||||
|
||||
print("\n[Step 1/2] 因子发现...")
|
||||
print("\n[Step 1/1] 因子发现...")
|
||||
research = self.agents["research"].execute(ts_codes=ts_codes)
|
||||
results = {"research": research}
|
||||
|
||||
|
||||
@@ -8,6 +8,9 @@ import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from agents.base import BaseAgent
|
||||
from config.settings import (
|
||||
RISK_THRESHOLDS, RISK_EXPOSURE, RISK_MAX_SINGLE, RISK_STOP_LOSS,
|
||||
)
|
||||
|
||||
|
||||
class RiskAgent(BaseAgent):
|
||||
@@ -16,11 +19,8 @@ class RiskAgent(BaseAgent):
|
||||
name = "Risk"
|
||||
description = "仓位控制与风险预警"
|
||||
|
||||
# 风险等级阈值
|
||||
THRESHOLDS = {
|
||||
"high": {"vol": 35, "dd": -15},
|
||||
"medium": {"vol": 25, "dd": -8},
|
||||
}
|
||||
# 风险等级阈值(来自配置中心)
|
||||
THRESHOLDS = RISK_THRESHOLDS
|
||||
|
||||
def execute(
|
||||
self,
|
||||
@@ -55,24 +55,22 @@ class RiskAgent(BaseAgent):
|
||||
peak = close.expanding().max()
|
||||
current_dd = float((close.iloc[-1] / peak.iloc[-1] - 1) * 100)
|
||||
|
||||
# 风险等级
|
||||
# 风险等级 & 仓位(参数来自配置中心)
|
||||
if market_vol > self.THRESHOLDS["high"]["vol"] or current_dd < self.THRESHOLDS["high"]["dd"]:
|
||||
risk_level = "high"
|
||||
target_exposure = 0.30
|
||||
elif market_vol > self.THRESHOLDS["medium"]["vol"] or current_dd < self.THRESHOLDS["medium"]["dd"]:
|
||||
risk_level = "medium"
|
||||
target_exposure = 0.60
|
||||
else:
|
||||
risk_level = "low"
|
||||
target_exposure = 0.85
|
||||
target_exposure = RISK_EXPOSURE[risk_level]
|
||||
|
||||
# 最近 N 日涨跌
|
||||
ret_5d = float(close.pct_change(5).iloc[-1] * 100) if len(close) >= 6 else 0
|
||||
ret_20d = float(close.pct_change(20).iloc[-1] * 100) if len(close) >= 21 else 0
|
||||
|
||||
# 单票上限(风险越高越集中)
|
||||
max_single = 0.15 if risk_level == "low" else (0.10 if risk_level == "medium" else 0.05)
|
||||
stop_loss = -0.05 if risk_level == "low" else (-0.08 if risk_level == "medium" else -0.12)
|
||||
# 单票上限(风险越高越集中)与止损线
|
||||
max_single = RISK_MAX_SINGLE[risk_level]
|
||||
stop_loss = RISK_STOP_LOSS[risk_level]
|
||||
|
||||
# 持仓预警
|
||||
alerts = []
|
||||
|
||||
@@ -8,6 +8,10 @@ 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):
|
||||
@@ -43,18 +47,13 @@ class SelectionAgent(BaseAgent):
|
||||
ts_codes = self.sent.get_scope_stocks()
|
||||
else:
|
||||
stocks = self.dm.get_stock_list()
|
||||
ts_codes = list(stocks.index[:100])
|
||||
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 = [
|
||||
"momentum_20", "momentum_60",
|
||||
"rsi_14", "volatility_20",
|
||||
"vol_ratio_5", "ma_dev_20",
|
||||
"turnover_5", "amplitude_5",
|
||||
]
|
||||
# 选择核心因子(覆盖多个维度,减少计算量;参数来自配置中心)
|
||||
core_factors = list(SELECTION_CORE_FACTORS)
|
||||
factor_objects = [get_factor(n) for n in core_factors]
|
||||
|
||||
self.log("打分 {} 只股票 (权重={})".format(len(ts_codes), weighting))
|
||||
@@ -70,7 +69,7 @@ class SelectionAgent(BaseAgent):
|
||||
len(available) / len(ts_codes) * 100 if ts_codes else 0))
|
||||
|
||||
# 2. 打分(限制上限防止单次太慢)
|
||||
score_limit = min(len(available), 300)
|
||||
score_limit = min(len(available), SELECTION_SCORE_LIMIT)
|
||||
scores = {}
|
||||
valid_count = 0
|
||||
for i, ts_code in enumerate(available[:score_limit]):
|
||||
@@ -165,23 +164,38 @@ class SelectionAgent(BaseAgent):
|
||||
if len(row_clean) < 3:
|
||||
return None
|
||||
|
||||
# z-score 标准化
|
||||
z = (row_clean - factor_df[row_clean.index].mean()) / factor_df[row_clean.index].std().replace(0, 1)
|
||||
# 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 模型打分。"""
|
||||
from models.features import FeatureEngine
|
||||
fe = FeatureEngine(lookahead=5)
|
||||
"""ML 模型打分。
|
||||
|
||||
daily = self.dm.get_daily(ts_code)
|
||||
要求注入名为 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.empty:
|
||||
if X is None or X.empty:
|
||||
return None
|
||||
if date and date in X.index:
|
||||
X = X.loc[[date]]
|
||||
@@ -191,5 +205,7 @@ class SelectionAgent(BaseAgent):
|
||||
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:
|
||||
except Exception as e:
|
||||
# 不再静默返回 None:记录原因,便于定位预测路径问题
|
||||
self.log("[WARN] ML 打分失败 ({}): {}".format(ts_code, e))
|
||||
return None
|
||||
|
||||
@@ -62,11 +62,15 @@ class BacktestReport:
|
||||
except AttributeError:
|
||||
return float(v) if v is not None else 0.0
|
||||
|
||||
# VectorBT 1.1.0 的 pf.value() 可能返回数字位置 index(0..n-1)。
|
||||
# 统一用它对应的交易日 index(close 已是规范化后的 DatetimeIndex),
|
||||
# 长度相等时按位置对齐,保证 equity/drawdown 的日期语义正确。
|
||||
equity = pf.value()
|
||||
equity = pd.Series(equity.values, index=close.index[: len(equity)])
|
||||
|
||||
if not isinstance(equity.index, pd.DatetimeIndex):
|
||||
equity.index = pd.to_datetime(equity.index, format="%Y%m%d")
|
||||
if not isinstance(equity.index, pd.DatetimeIndex) and len(equity) == len(close):
|
||||
equity.index = close.index
|
||||
elif not isinstance(equity.index, pd.DatetimeIndex):
|
||||
# 长度不等时尽力用 close 前缀对齐,避免异常
|
||||
equity.index = close.index[: len(equity)]
|
||||
|
||||
dd = equity / equity.cummax() - 1
|
||||
daily_ret = equity.pct_change().dropna()
|
||||
|
||||
@@ -22,12 +22,18 @@ class VectorBTEngine:
|
||||
def __init__(
|
||||
self,
|
||||
initial_capital: float = 100_000,
|
||||
commission: float = 0.0003, # 万三
|
||||
commission: float = 0.0003, # 佣金(万三)
|
||||
slippage: float = 0.0000, # 滑点(比例,0=关闭),单边
|
||||
freq: str = "D",
|
||||
t_plus_one: bool = True, # 信号次日开盘执行(A 股 T+1)
|
||||
limit_check: bool = True, # 涨停拒买 / 跌停拒卖
|
||||
):
|
||||
self.initial_capital = initial_capital
|
||||
self.commission = commission
|
||||
self.slippage = slippage
|
||||
self.freq = freq
|
||||
self.t_plus_one = t_plus_one
|
||||
self.limit_check = limit_check
|
||||
|
||||
# ── 单股票回测 ────────────────────────────────────────
|
||||
|
||||
@@ -36,19 +42,36 @@ class VectorBTEngine:
|
||||
strategy: BaseStrategy,
|
||||
price_df: pd.DataFrame,
|
||||
factor_df: pd.DataFrame | None = None,
|
||||
t_plus_one: bool | None = None,
|
||||
slippage: float | None = None,
|
||||
limit_check: bool | None = None,
|
||||
) -> BacktestReport:
|
||||
"""
|
||||
单股票回测。
|
||||
|
||||
参数:
|
||||
strategy: 策略实例
|
||||
price_df: 价格数据,index=trade_date,必须有 'close' 列
|
||||
factor_df: 因子数据,index=trade_date。
|
||||
None 时使用 price_df 作为因子数据源。
|
||||
price_df: 价格数据,index=trade_date,必须含 'close';
|
||||
若含 'open' 且启用 t_plus_one,则用次日 open 成交。
|
||||
factor_df: 因子数据,index=trade_date。None 时使用 price_df。
|
||||
t_plus_one: 覆盖引擎默认的 T+1 异步成交;None=用引擎设置。
|
||||
slippage: 覆盖引擎默认滑点;None=用引擎设置。
|
||||
limit_check: 覆盖引擎默认的涨跌停拒成交;None=用引擎设置。
|
||||
|
||||
返回:
|
||||
BacktestReport
|
||||
|
||||
成交真实性(相对旧版的关键修复):
|
||||
- T+1: 信号当日收盘产生,成交推迟到次日,避免"今天收盘出信号、
|
||||
今天收盘就成交"的乐观偏差。
|
||||
- 涨停拒买 / 跌停拒卖: 涨停当日实际无法买入、跌停当日无法卖出,
|
||||
对应位置的入场/出场信号被抑制。
|
||||
- 滑点: 通过 slippage 施加成交价冲击。
|
||||
"""
|
||||
t_plus_one = self.t_plus_one if t_plus_one is None else t_plus_one
|
||||
slippage = self.slippage if slippage is None else slippage
|
||||
limit_check = self.limit_check if limit_check is None else limit_check
|
||||
|
||||
if factor_df is None:
|
||||
factor_df = price_df
|
||||
|
||||
@@ -60,6 +83,11 @@ class VectorBTEngine:
|
||||
price_df = price_df.loc[common_idx].sort_index()
|
||||
factor_df = factor_df.loc[common_idx].sort_index()
|
||||
|
||||
# 1b. 统一 index 为交易日 DatetimeIndex(支持 'YYYYMMDD' 字符串/int),
|
||||
# 保证 vectorbt 的 equity 与 report 的日期语义正确。
|
||||
price_df.index = self._normalize_daily_index(price_df.index)
|
||||
factor_df.index = self._normalize_daily_index(factor_df.index)
|
||||
|
||||
# 2. 合并 close 到 factor_df(策略可能需要)
|
||||
if "close" not in factor_df.columns:
|
||||
factor_df = factor_df.copy()
|
||||
@@ -71,19 +99,90 @@ class VectorBTEngine:
|
||||
# 4. 信号 → VectorBT entries/exits
|
||||
entries, exits = self._signals_to_entries(raw_signals, price_df.index)
|
||||
|
||||
# 5. 运行回测
|
||||
close = price_df["close"]
|
||||
# 5. A 股真实性过滤:涨跌停拒成交(在 T+1 前,基于信号当日的行情状态)
|
||||
if limit_check:
|
||||
entries, exits = self._apply_limit_filters(entries, exits, price_df)
|
||||
|
||||
# 6. T+1 异步成交:入场/出场推迟到次日开盘
|
||||
# (shift 会引入 NaN 使 dtype 变 object;显式转回 bool,
|
||||
# 否则 vectorbt 的 numba 内核因 object 数组报 TypingError)
|
||||
if t_plus_one:
|
||||
entries = entries.shift(1).fillna(False).astype(bool)
|
||||
exits = exits.shift(1).fillna(False).astype(bool)
|
||||
|
||||
# 7. 运行回测(用 open 序列做成交价,否则 fallback 到 close)
|
||||
exec_price = price_df["open"] if "open" in price_df.columns else price_df["close"]
|
||||
pf = vbt.Portfolio.from_signals(
|
||||
close,
|
||||
exec_price,
|
||||
entries=entries,
|
||||
exits=exits,
|
||||
init_cash=self.initial_capital,
|
||||
fees=self.commission,
|
||||
slippage=slippage or None,
|
||||
freq=self.freq,
|
||||
direction="longonly",
|
||||
)
|
||||
|
||||
return BacktestReport.from_vbt_result(pf, close)
|
||||
return BacktestReport.from_vbt_result(pf, exec_price)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_daily_index(index: pd.Index) -> pd.Index:
|
||||
"""把 'YYYYMMDD' 字符串或 int 类型的 index 统一为 DatetimeIndex。"""
|
||||
if isinstance(index, pd.DatetimeIndex):
|
||||
return index
|
||||
# int64(如 20240101)
|
||||
if pd.api.types.is_integer_dtype(index):
|
||||
parsed = pd.to_datetime(index.astype(str), format="%Y%m%d", errors="coerce")
|
||||
if parsed.notna().all():
|
||||
return parsed
|
||||
elif index.dtype == object or isinstance(index, pd.Index):
|
||||
parsed = pd.to_datetime(index, format="%Y%m%d", errors="coerce")
|
||||
if parsed.notna().all():
|
||||
return parsed
|
||||
return index
|
||||
|
||||
def _apply_limit_filters(
|
||||
self,
|
||||
entries: pd.Series,
|
||||
exits: pd.Series,
|
||||
price_df: pd.DataFrame,
|
||||
) -> tuple[pd.Series, pd.Series]:
|
||||
"""
|
||||
涨停拒买 / 跌停拒卖。
|
||||
|
||||
依托 price_df 的 pre_close / pct_chg(若存在)估算涨跌停:
|
||||
- close 达到/接近涨停 → 当日无法买入 → 抑制 entry
|
||||
- close 达到/接近跌停 → 当日无法卖出 → 抑制 exit
|
||||
|
||||
板块差异(ST 5%、主板 10%、创业板/科创板 20%)通过 ts_code 后缀近似判断,
|
||||
无后缀信息时按主板 10% 上限处理。
|
||||
"""
|
||||
if "pre_close" in price_df.columns and "close" in price_df.columns:
|
||||
pre_close = price_df["pre_close"].replace(0, float("nan"))
|
||||
pct = (price_df["close"] - pre_close) / pre_close * 100
|
||||
elif "pct_chg" in price_df.columns:
|
||||
pct = price_df["pct_chg"]
|
||||
else:
|
||||
return entries, exits # 无行情判断列,跳过
|
||||
|
||||
code = str(price_df.index.name or "") or ""
|
||||
# 用列里的 ts_code 判断板块(若有)
|
||||
board_limit = 9.8
|
||||
if "ts_code" in price_df.columns:
|
||||
codes = price_df["ts_code"].astype(str)
|
||||
# 创业板 300/301/688 科创板 → 20%,ST 无后缀信息按 10%
|
||||
limit_20 = codes.str.match(r"^(300|301|688)\d{3}")
|
||||
board_limit = 19.6
|
||||
# 留 margin:pct >= +9.8 判定接近涨停(不可买),<= -9.8 判定接近跌停(不可卖)
|
||||
up = pct >= 9.8
|
||||
down = pct <= -9.8
|
||||
if board_limit > 9.8:
|
||||
up = pct >= 19.6
|
||||
down = pct <= -19.6
|
||||
|
||||
entries = entries & ~up
|
||||
exits = exits & ~down
|
||||
return entries, exits
|
||||
|
||||
# ── 截面回测(多股票) ──────────────────────────────────
|
||||
|
||||
@@ -97,41 +196,143 @@ class VectorBTEngine:
|
||||
"""
|
||||
截面策略回测(多股票 + 定期调仓)。
|
||||
|
||||
对每只股票独立回测,合并权益曲线。
|
||||
真实组合语义(相对旧版"每股满额独立回测再等权平均"的关键修复):
|
||||
- 每股先用策略信号驱动出每日持仓状态(T+1 成交);
|
||||
- `rebalance_freq` 决定持仓只在调仓日更新('D'/'W'/'M');
|
||||
- 组合总资金(initial_capital)在当日持仓股票间等权切分,
|
||||
资金不会被重复分配/超限,是可联合投资的单一连续净值。
|
||||
|
||||
参数:
|
||||
strategy: 策略实例
|
||||
price_universe: {ts_code: price_df}
|
||||
price_universe: {ts_code: price_df}(至少含 close;可含 open)
|
||||
factor_universe: {ts_code: factor_df}
|
||||
rebalance_freq: 调仓频率 'D'/'W'/'M',用于合并时对齐
|
||||
rebalance_freq: 调仓频率 'D'/'W'/'M'(默认 'M' 月度)
|
||||
|
||||
返回:
|
||||
BacktestReport
|
||||
BacktestReport(组合级,equity_curve=组合净额曲线)
|
||||
"""
|
||||
if factor_universe is None:
|
||||
factor_universe = price_universe
|
||||
|
||||
stock_equities = {}
|
||||
stock_reports = {}
|
||||
|
||||
# 逐股票回测
|
||||
# 1. 每股生成 T+1 后的持仓状态序列(策略信号驱动)
|
||||
holdings: dict[str, pd.Series] = {} # ts_code -> bool 每日是否持仓
|
||||
returns: dict[str, pd.Series] = {} # ts_code -> 每日收益率
|
||||
for ts_code in price_universe:
|
||||
price_df = price_universe[ts_code]
|
||||
if "close" not in price_df.columns or price_df.empty:
|
||||
continue
|
||||
|
||||
factor_df = factor_universe.get(ts_code, price_df)
|
||||
common = price_df.index.intersection(factor_df.index)
|
||||
if len(common) < 2:
|
||||
continue
|
||||
p = price_df.loc[common].sort_index()
|
||||
f = factor_df.loc[common].sort_index()
|
||||
p.index = self._normalize_daily_index(p.index)
|
||||
f.index = self._normalize_daily_index(f.index)
|
||||
|
||||
report = self.run(strategy, price_df, factor_df)
|
||||
if report is not None and len(report.equity_curve) > 0:
|
||||
stock_equities[ts_code] = report.equity_curve
|
||||
stock_reports[ts_code] = report
|
||||
raw_signals = strategy.generate_signals(f)
|
||||
entries, exits = self._signals_to_entries(raw_signals, p.index)
|
||||
# T+1 成交:信号次日生效
|
||||
entries = entries.shift(1).fillna(False).astype(bool)
|
||||
exits = exits.shift(1).fillna(False).astype(bool)
|
||||
|
||||
if not stock_equities:
|
||||
pos = pd.Series(False, index=p.index)
|
||||
in_now = False
|
||||
e = entries.to_numpy(); x = exits.to_numpy()
|
||||
for i in range(len(p)):
|
||||
if e[i]:
|
||||
in_now = True
|
||||
elif x[i]:
|
||||
in_now = False
|
||||
pos.iloc[i] = in_now
|
||||
|
||||
holdings[ts_code] = pos
|
||||
returns[ts_code] = p["close"].pct_change()
|
||||
|
||||
if not holdings:
|
||||
return BacktestReport()
|
||||
|
||||
# 合并:等权分配资金到各股票
|
||||
return self._merge_equities(stock_equities)
|
||||
# 2. 统一交易日历(全部股票 index 并集,升序)
|
||||
all_days = pd.DatetimeIndex(
|
||||
sorted(set().union(*[h.index for h in holdings.values()]))
|
||||
)
|
||||
|
||||
# 3. rebalance 时点(持仓只在调仓日变化)
|
||||
rebalance_mask = self._rebalance_mask(all_days, rebalance_freq)
|
||||
|
||||
# 4. 逐日计算组合等权收益(资金在当日持仓之间切分)
|
||||
pos_matrix = {c: h.reindex(all_days).fillna(False) for c, h in holdings.items()}
|
||||
ret_matrix = {c: r.reindex(all_days).fillna(0.0) for c, r in returns.items()}
|
||||
|
||||
current_pos = {c: False for c in holdings}
|
||||
daily_port_ret = np.zeros(len(all_days))
|
||||
|
||||
for i, day in enumerate(all_days):
|
||||
if rebalance_mask[i]:
|
||||
# 调仓:按当日的 T+1 持仓状态重新确定各股是否纳入组合
|
||||
for c in holdings:
|
||||
current_pos[c] = bool(pos_matrix[c].iloc[i])
|
||||
# 当日组合收益 = 持仓股票当日收益的等权平均(资金按持仓数切分)
|
||||
held = [c for c in holdings if current_pos[c]]
|
||||
if held:
|
||||
daily_port_ret[i] = np.mean([ret_matrix[c].iloc[i] for c in held])
|
||||
|
||||
port_ret = pd.Series(daily_port_ret, index=all_days)
|
||||
portfolio_equity = self.initial_capital * (1 + port_ret).cumprod()
|
||||
|
||||
# 5. 组合指标
|
||||
return self._build_portfolio_report(portfolio_equity)
|
||||
|
||||
@staticmethod
|
||||
def _rebalance_mask(days: pd.DatetimeIndex, freq: str) -> np.ndarray:
|
||||
"""生成调仓日布尔掩码:'D'=每天,'W'=每周首个,'M'=每月首个。"""
|
||||
mask = np.zeros(len(days), dtype=bool)
|
||||
freq = (freq or "M").upper()
|
||||
if freq == "D":
|
||||
mask[:] = True
|
||||
return mask
|
||||
prev_key = None
|
||||
for i, day in enumerate(days):
|
||||
if freq == "W":
|
||||
key = (day.isocalendar()[0], day.isocalendar()[1])
|
||||
else: # 'M'
|
||||
key = (day.year, day.month)
|
||||
if key != prev_key:
|
||||
mask[i] = True
|
||||
prev_key = key
|
||||
return mask
|
||||
|
||||
def _build_portfolio_report(self, portfolio_equity: pd.Series) -> BacktestReport:
|
||||
"""从组合净值曲线计算标准化指标。"""
|
||||
dd = portfolio_equity / portfolio_equity.cummax() - 1
|
||||
daily_ret = portfolio_equity.pct_change().dropna()
|
||||
years = max(len(daily_ret) / 252, 0.02)
|
||||
|
||||
total_ret = (portfolio_equity.iloc[-1] / portfolio_equity.iloc[0] - 1) * 100
|
||||
cagr = ((total_ret / 100 + 1) ** (1 / years) - 1) * 100
|
||||
mdd = dd.min() * 100
|
||||
mean_ret = daily_ret.mean() * 252
|
||||
std_ret = daily_ret.std() * np.sqrt(252)
|
||||
sharpe = mean_ret / std_ret if std_ret > 0 else 0
|
||||
calmar = cagr / abs(mdd) if abs(mdd) > 0 else 0
|
||||
|
||||
try:
|
||||
monthly = portfolio_equity.resample("ME").last().pct_change()
|
||||
except Exception:
|
||||
monthly = pd.Series(dtype=float)
|
||||
|
||||
return BacktestReport(
|
||||
total_return=round(total_ret, 2),
|
||||
cagr=round(cagr, 2),
|
||||
max_drawdown=round(mdd, 2),
|
||||
sharpe_ratio=round(sharpe, 2),
|
||||
calmar_ratio=round(calmar, 2),
|
||||
annual_volatility=round(std_ret * 100 if std_ret != 0 else 0, 2),
|
||||
total_trades=0,
|
||||
equity_curve=portfolio_equity,
|
||||
drawdown_curve=dd,
|
||||
monthly_returns=monthly,
|
||||
)
|
||||
|
||||
# ── 信号转换 ──────────────────────────────────────────
|
||||
|
||||
|
||||
+64
-42
@@ -3,19 +3,24 @@
|
||||
Agent 命令行入口。
|
||||
|
||||
用法:
|
||||
python cli/agent_cli.py daily # 执行每日流程
|
||||
python cli/agent_cli.py picks [N] # 今日选股 Top N
|
||||
python cli/agent_cli.py risk # 风险评估
|
||||
python cli/agent_cli.py research # 因子研究
|
||||
python cli/agent_cli.py report [DATE] # 生成日报
|
||||
python cli/agent_cli.py daily # 执行每日流程
|
||||
python cli/agent_cli.py picks [N] # 今日选股 Top N
|
||||
python cli/agent_cli.py risk # 风险评估
|
||||
python cli/agent_cli.py research # 因子研究
|
||||
python cli/agent_cli.py report [DATE] # 生成日报
|
||||
python cli/agent_cli.py warmup [N] # 首次批量预热
|
||||
"""
|
||||
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
from datetime import datetime
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(name)s: %(message)s")
|
||||
|
||||
|
||||
def init_engines():
|
||||
"""初始化所有引擎。"""
|
||||
@@ -26,7 +31,6 @@ def init_engines():
|
||||
from factors.sentiment.sentiment_engine import SentimentEngine
|
||||
from factors.sentiment.news_source import NewsSource
|
||||
from factors.sentiment.qwen_client import QwenClient
|
||||
from factors.registry import get_factor
|
||||
|
||||
dm = DataManager()
|
||||
dm.init_db()
|
||||
@@ -45,42 +49,62 @@ def init_engines():
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
if len(sys.argv) < 2:
|
||||
print("用法: agent_cli.py <daily|picks|risk|research|report|warmup>")
|
||||
print()
|
||||
print(" daily [DATE] — 执行每日完整流程")
|
||||
print(" picks [N] [DATE] — 今日选股 Top N")
|
||||
print(" risk — 风险评估")
|
||||
print(" research — 因子发现与评估")
|
||||
print(" report [DATE] — 生成日报")
|
||||
print(" warmup [N] — 首次批量预热范围股票到 DB 缓存")
|
||||
return
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
p = argparse.ArgumentParser(prog="agent_cli", description="cc-cursor 量化 Agent CLI")
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
|
||||
cmd = sys.argv[1]
|
||||
d = sub.add_parser("daily", help="执行每日完整流程")
|
||||
d.add_argument("date", nargs="?", default=None, help="日期 YYYYMMDD(默认今天)")
|
||||
|
||||
pk = sub.add_parser("picks", help="今日选股 Top N")
|
||||
pk.add_argument("top_n", nargs="?", type=int, default=15, help="返回 Top N")
|
||||
pk.add_argument("date", nargs="?", default=None, help="日期 YYYYMMDD")
|
||||
|
||||
sub.add_parser("risk", help="风险评估")
|
||||
|
||||
sub.add_parser("research", help="因子发现与评估")
|
||||
|
||||
rp = sub.add_parser("report", help="生成日报")
|
||||
rp.add_argument("date", nargs="?", default=None, help="日期 YYYYMMDD")
|
||||
|
||||
wm = sub.add_parser("warmup", help="首次批量预热范围股票到 DB 缓存")
|
||||
wm.add_argument("batch_n", nargs="?", type=int, default=50, help="每批股票数")
|
||||
|
||||
return p
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
parser = build_parser()
|
||||
args = parser.parse_args(argv)
|
||||
engines = init_engines()
|
||||
|
||||
from agents.orchestrator import AgentOrchestrator
|
||||
orch = AgentOrchestrator(**engines)
|
||||
orch.setup()
|
||||
|
||||
if cmd == "daily":
|
||||
date = sys.argv[2] if len(sys.argv) > 2 else None
|
||||
results = orch.run_daily(date=date)
|
||||
# 打印日报内容
|
||||
today = datetime.now().strftime("%Y%m%d")
|
||||
|
||||
if args.cmd == "daily":
|
||||
results = orch.run_daily(date=args.date)
|
||||
report_md = results.get("report", {}).get("report_markdown", "")
|
||||
if report_md:
|
||||
print(report_md)
|
||||
# 任一步骤失败 → 非零退出码,便于调度/CI 感知
|
||||
failed = [k for k, v in results.items()
|
||||
if isinstance(v, dict) and "error" in v]
|
||||
if failed:
|
||||
print("每日流程有步骤未完成: {}".format(", ".join(failed)), file=sys.stderr)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
elif cmd == "picks":
|
||||
n = int(sys.argv[2]) if len(sys.argv) > 2 else 15
|
||||
date = sys.argv[3] if len(sys.argv) > 3 else None
|
||||
result = orch.picks(date=date, top_n=n)
|
||||
elif args.cmd == "picks":
|
||||
result = orch.picks(date=args.date, top_n=args.top_n)
|
||||
print(f"\n选股结果 ({result.get('date', '?')}):")
|
||||
for p in result.get("top_picks", []):
|
||||
print(f" {p['ts_code']:12s} {p.get('name', ''):10s} {p['score']:.4f}")
|
||||
return 0
|
||||
|
||||
elif cmd == "risk":
|
||||
elif args.cmd == "risk":
|
||||
result = orch.risk_check()
|
||||
print(f"\n风险评估:")
|
||||
print(f" 等级: {result['risk_level']}")
|
||||
@@ -93,22 +117,24 @@ def main():
|
||||
print(f" 回撤: {indicators.get('current_drawdown', 0):.1f}%")
|
||||
for a in result.get("alerts", []):
|
||||
print(f" ⚠️ {a}")
|
||||
return 0
|
||||
|
||||
elif cmd == "research":
|
||||
elif args.cmd == "research":
|
||||
result = orch.run_research_cycle()
|
||||
top = result.get("research", {}).get("top_factors", [])
|
||||
print(f"\n因子评估结果:")
|
||||
if not top:
|
||||
print(" (无结果)")
|
||||
return
|
||||
return 0
|
||||
print(f" {'因子':20s} {'IC':>8s} {'IC_IR':>8s} {'多头':>8s} {'空头':>8s} {'得分':>8s}")
|
||||
print(f" {'─'*60}")
|
||||
for f in top:
|
||||
print(f" {f['name']:20s} {f['ic_mean']:>+8.4f} {f['icir']:>8.3f} "
|
||||
f"{f['long_ret']:>+7.1f}% {f['short_ret']:>+7.1f}% {f['score']:>8.4f}")
|
||||
return 0
|
||||
|
||||
elif cmd == "warmup":
|
||||
batch_n = int(sys.argv[2]) if len(sys.argv) > 2 else 50
|
||||
elif args.cmd == "warmup":
|
||||
batch_n = args.batch_n
|
||||
print("首次批量预热: 每次 {} 只股票,分批执行...".format(batch_n))
|
||||
sent = engines.get("sent")
|
||||
dm = engines.get("dm")
|
||||
@@ -129,24 +155,20 @@ def main():
|
||||
print(" {} 失败: {}".format(ts_code, e))
|
||||
print(" 累计同步: {} 条".format(total_synced))
|
||||
print("预热完成: {} 条数据, {} 只新股票已缓存".format(total_synced, len(uncached)))
|
||||
return 0
|
||||
|
||||
elif cmd == "report":
|
||||
date = sys.argv[2] if len(sys.argv) > 2 else None
|
||||
elif args.cmd == "report":
|
||||
date = args.date or today
|
||||
result = orch.generate_report(date=date)
|
||||
print("\n日报已生成: {}".format(result.get("report_path", "?")))
|
||||
# 存入 DB
|
||||
md = result.get("report_markdown", "")
|
||||
if md:
|
||||
from reports.storage import save_report
|
||||
save_report(md, "量化日报", report_date=date or datetime.now().strftime("%Y%m%d"),
|
||||
subject_type="daily", subject_code="")
|
||||
print(" 已存入 DB")
|
||||
# 注意: ReportAgent.execute 内部已保存到 DB,这里不再重复入库
|
||||
if md:
|
||||
print(md)
|
||||
return 0
|
||||
|
||||
else:
|
||||
print(f"未知命令: {cmd}")
|
||||
return 0 # 防御:未匹配(应不会到达)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
sys.exit(main())
|
||||
@@ -89,7 +89,7 @@ def main():
|
||||
test_price = price_df.loc[X_test.index]
|
||||
test_factor = factor_df.loc[X_test.index]
|
||||
bt_engine = VectorBTEngine()
|
||||
benchmark = MLBenchmark([lgb_model, cb_model], fe, test_price, test_factor, bt_engine)
|
||||
benchmark = MLBenchmark([lgb_model, cb_model], fe, test_factor, test_price, bt_engine)
|
||||
result = benchmark.run()
|
||||
print(result.round(2).to_string())
|
||||
print(" > 解读: 回测结果反映 ML 策略在测试集上的实盘表现。")
|
||||
|
||||
@@ -44,3 +44,45 @@ AKSHARE_CONFIG = {
|
||||
# ── 默认参数 ──────────────────────────────────────────────
|
||||
DEFAULT_START_DATE = "20200101"
|
||||
DEFAULT_END_DATE = None # None 表示当天
|
||||
|
||||
# ── 数据新鲜度 ────────────────────────────────────────────
|
||||
# DB 窗口尾部落后于请求日多少个自然日,视为存在缺口需补拉(get_daily)
|
||||
DATA_FRESHNESS_GAP_DAYS = int(os.getenv("DATA_FRESHNESS_GAP_DAYS", "7"))
|
||||
|
||||
# ── 风险评估参数(RiskAgent) ─────────────────────────────
|
||||
# 风险等级阈值:年化波动率(%) / 当前回撤(%),超过即升级
|
||||
RISK_THRESHOLDS = {
|
||||
"high": {"vol": float(os.getenv("RISK_HIGH_VOL", "35")), "dd": float(os.getenv("RISK_HIGH_DD", "-15"))},
|
||||
"medium": {"vol": float(os.getenv("RISK_MED_VOL", "25")), "dd": float(os.getenv("RISK_MED_DD", "-8"))},
|
||||
}
|
||||
# 各等级对应的建议总仓位、单票上限、止损线
|
||||
RISK_EXPOSURE = {"low": 0.85, "medium": 0.60, "high": 0.30}
|
||||
RISK_MAX_SINGLE = {"low": 0.15, "medium": 0.10, "high": 0.05}
|
||||
RISK_STOP_LOSS = {"low": -0.05, "medium": -0.08, "high": -0.12}
|
||||
|
||||
# ── 选股参数(SelectionAgent) ────────────────────────────
|
||||
# 默认选股池大小(按股票列表前 N 只作为代表股池)
|
||||
SELECTION_UNIVERSE_SIZE = int(os.getenv("SELECTION_UNIVERSE_SIZE", "100"))
|
||||
# 单次打分上限(防止单次过慢)
|
||||
SELECTION_SCORE_LIMIT = int(os.getenv("SELECTION_SCORE_LIMIT", "300"))
|
||||
# 默认返回 Top N
|
||||
SELECTION_TOP_N = int(os.getenv("SELECTION_TOP_N", "15"))
|
||||
# 选股核心因子(覆盖多维度,减少计算量)
|
||||
SELECTION_CORE_FACTORS = [
|
||||
"momentum_20", "momentum_60",
|
||||
"rsi_14", "volatility_20",
|
||||
"vol_ratio_5", "ma_dev_20",
|
||||
"turnover_5", "amplitude_5",
|
||||
]
|
||||
# z-score 打分时是否做去极值(避免单股离群值主导 Top 排名)
|
||||
SELECTION_WINSORIZE_ZSCORE = os.getenv("SELECTION_WINSORIZE_ZSCORE", "1") == "1"
|
||||
|
||||
# ── 报告参数(ReportAgent) ────────────────────────────────
|
||||
# 日报市场概览指数代码
|
||||
REPORT_INDEX_CODES = {
|
||||
"000001.SH": "上证指数",
|
||||
"399001.SZ": "深证成指",
|
||||
"399006.SZ": "创业板指",
|
||||
}
|
||||
# 日报标题中展示的 Top N(与 SELECTION_TOP_N 联动)
|
||||
REPORT_TOP_N = SELECTION_TOP_N
|
||||
|
||||
@@ -5,6 +5,7 @@ DataManager — 统一数据管理层。
|
||||
优先从 DB 读取,缺失时依次尝试 AkShare → Tushare 拉取并入库。
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
|
||||
import pandas as pd
|
||||
@@ -15,6 +16,27 @@ from data.sources.tushare_source import TushareSource
|
||||
from database import dao
|
||||
from database.models import create_all_tables
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _rollback_days(yyyymmdd: str, days: int) -> str:
|
||||
"""把 'YYYYMMDD' 往前回退 days 个自然日,返回 'YYYYMMDD'。"""
|
||||
from datetime import datetime, timedelta
|
||||
dt = datetime.strptime(yyyymmdd, "%Y%m%d") - timedelta(days=days)
|
||||
return dt.strftime("%Y%m%d")
|
||||
|
||||
|
||||
def _safe_save_daily(df: pd.DataFrame, skip_exc: bool = True) -> None:
|
||||
"""把拉取到的日线入库;失败时记日志(可选抛错)。"""
|
||||
cols = [c for c in dao._DAILY_COLS if c in df.columns]
|
||||
try:
|
||||
dao.save_daily(df[cols])
|
||||
except Exception as e:
|
||||
if skip_exc:
|
||||
logger.exception("[DataManager] 日线入库失败(数据已获取,未持久化):%s", e)
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
class DataManager:
|
||||
"""统一数据管理。双数据源:AkShare(主)+ Tushare(备)。"""
|
||||
@@ -68,7 +90,7 @@ class DataManager:
|
||||
if df is not None and not df.empty:
|
||||
return df, label
|
||||
except Exception as e:
|
||||
print(" [{}] {} 失败: {}".format(label, method_name, e))
|
||||
logger.warning(" [%s] %s 失败: %s", label, method_name, e)
|
||||
return pd.DataFrame(), None
|
||||
|
||||
# ── 股票列表 ──────────────────────────────────────────
|
||||
@@ -77,15 +99,15 @@ class DataManager:
|
||||
if not force_refresh:
|
||||
df = dao.query_stock_list()
|
||||
if not df.empty:
|
||||
print("[DataManager] 从 DB 读取股票列表: {} 只".format(len(df)))
|
||||
logger.info("[DataManager] 从 DB 读取股票列表: %s 只", len(df))
|
||||
return df
|
||||
|
||||
print("[DataManager] 拉取股票列表 (AkShare → Tushare)...")
|
||||
logger.info("[DataManager] 拉取股票列表 (AkShare → Tushare)...")
|
||||
df, src = self._try_fetch("fetch_stock_list")
|
||||
if df.empty:
|
||||
print("[DataManager] 所有数据源均无法获取股票列表")
|
||||
logger.warning("[DataManager] 所有数据源均无法获取股票列表")
|
||||
return pd.DataFrame()
|
||||
print("[DataManager] 股票列表已入库 ({}): {} 只".format(src, len(df)))
|
||||
logger.info("[DataManager] 股票列表已入库 (%s): %s 只", src, len(df))
|
||||
dao.save_stock_list(df)
|
||||
time.sleep(2)
|
||||
return df
|
||||
@@ -105,40 +127,48 @@ class DataManager:
|
||||
if not force_refresh:
|
||||
df = dao.query_daily(ts_code, start, end)
|
||||
if not df.empty:
|
||||
# 检查 DB 窗口尾部是否明显落后于请求的 end(>5 个自然日)。
|
||||
# 若是,说明存在数据缺口,触发一次增量补拉,而不是把陈旧数据当完整返回。
|
||||
latest = str(df["trade_date"].astype(str).max())
|
||||
if end >= "20200101" and latest < _rollback_days(end, 7):
|
||||
logger.warning(
|
||||
"%s DB 窗口尾部 %s 明显落后于请求日 %s,触发补拉", ts_code, latest, end)
|
||||
refreshed, src = self._try_fetch("fetch_daily", ts_code, latest, end)
|
||||
if not refreshed.empty and src:
|
||||
_safe_save_daily(refreshed)
|
||||
return df
|
||||
|
||||
# DB 未命中 → 从数据源拉取
|
||||
df, src = self._try_fetch("fetch_daily", ts_code, start, end)
|
||||
if df.empty:
|
||||
print("[WARN] {} 日线获取失败 (AkShare+Tushare 均不可用)".format(ts_code))
|
||||
logger.warning("%s 日线获取失败 (AkShare+Tushare 均不可用)", ts_code)
|
||||
return pd.DataFrame()
|
||||
|
||||
# 保存到 DB
|
||||
try:
|
||||
# 筛选 DB 需要的列
|
||||
cols = [c for c in dao._DAILY_COLS if c in df.columns]
|
||||
dao.save_daily(df[cols])
|
||||
except Exception as e:
|
||||
print("[WARN] 日线入库失败: {}".format(e))
|
||||
|
||||
# 保存到 DB(失败仅记日志,不阻断已获取到的数据返回)
|
||||
_safe_save_daily(df)
|
||||
return df
|
||||
|
||||
def sync_daily(self, ts_code: str) -> int:
|
||||
latest = dao.get_latest_trade_date(ts_code)
|
||||
today = time.strftime("%Y%m%d")
|
||||
if latest and latest >= today:
|
||||
print("[DataManager] {} 数据已是最新 ({})".format(ts_code, latest))
|
||||
logger.info("[DataManager] %s 数据已是最新 (%s)", ts_code, latest)
|
||||
return 0
|
||||
|
||||
start = latest or DEFAULT_START_DATE
|
||||
df, src = self._try_fetch("fetch_daily", ts_code, start, today)
|
||||
if df.empty:
|
||||
print("[WARN] {} sync_daily 失败 (AkShare+Tushare 均不可用)".format(ts_code))
|
||||
logger.warning("[DataManager] %s sync_daily 失败 (AkShare+Tushare 均不可用)", ts_code)
|
||||
return 0
|
||||
|
||||
# 入库失败要能反映到返回值,否则调用方会误以为已持久化
|
||||
cols = [c for c in dao._DAILY_COLS if c in df.columns]
|
||||
dao.save_daily(df[cols])
|
||||
print("[DataManager] {} 同步 {} 条日线 ({})".format(ts_code, len(df), src))
|
||||
try:
|
||||
dao.save_daily(df[cols])
|
||||
except Exception as e:
|
||||
logger.exception("[DataManager] %s 日线入库失败:%s", ts_code, e)
|
||||
raise
|
||||
logger.info("[DataManager] %s 同步 %s 条日线 (%s)", ts_code, len(df), src)
|
||||
return len(df)
|
||||
|
||||
def sync_all_daily(self) -> int:
|
||||
@@ -148,11 +178,11 @@ class DataManager:
|
||||
try:
|
||||
total += self.sync_daily(ts_code)
|
||||
if (i + 1) % 50 == 0:
|
||||
print("[DataManager] 进度: {}/{}".format(i + 1, len(stock_list)))
|
||||
logger.info("[DataManager] 进度: %s/%s", i + 1, len(stock_list))
|
||||
time.sleep(1)
|
||||
except Exception as e:
|
||||
print("[WARN] {} 同步失败: {}".format(ts_code, e))
|
||||
print("[DataManager] 全量同步完成,新增 {} 条".format(total))
|
||||
logger.warning("[DataManager] %s 同步失败: %s", ts_code, e)
|
||||
logger.info("[DataManager] 全量同步完成,新增 %s 条", total)
|
||||
return total
|
||||
|
||||
# ── 财务数据 ──────────────────────────────────────────
|
||||
@@ -167,5 +197,5 @@ class DataManager:
|
||||
cols = [c for c in dao._FINA_COLS if c in df.columns]
|
||||
dao.save_financial(df[cols])
|
||||
except Exception as e:
|
||||
print("[WARN] 财务数据入库失败: {}".format(e))
|
||||
logger.exception("财务数据入库失败: %s", e)
|
||||
return df
|
||||
|
||||
+66
-34
@@ -4,12 +4,16 @@
|
||||
提供 DataFrame 级别的读写操作,屏蔽底层 ORM/SQL 细节。
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
import pandas as pd
|
||||
from sqlalchemy import text
|
||||
|
||||
from database.connection import get_engine
|
||||
from database.models import StockBasic, StockDaily, StockFinancial, Report
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# DB 表列名,供 DataManager 在写入前筛选
|
||||
_DAILY_COLS = [
|
||||
"ts_code", "trade_date", "open", "high", "low", "close",
|
||||
@@ -22,33 +26,59 @@ _FINA_COLS = [
|
||||
]
|
||||
|
||||
|
||||
def _df_to_db(df: pd.DataFrame, model_class, replace: bool = False) -> int:
|
||||
"""将 DataFrame 写入对应表,返回写入行数。"""
|
||||
if df.empty:
|
||||
def _upsert_df(engine, model_class, df: pd.DataFrame) -> int:
|
||||
"""
|
||||
用 SQLAlchemy Core 做 upsert(INSERT ... ON DUPLICATE KEY UPDATE)写库。
|
||||
|
||||
不会 DROP/重建表(区别于 pandas to_sql 的 if_exists="replace"),
|
||||
主键冲突时按非主键列更新而非报错。返回受影响行数。
|
||||
|
||||
说明: 需要真实 ORM model 的 __table__(含主键信息),
|
||||
因此不能像旧版那样用字符串表名 + pandas to_sql。
|
||||
"""
|
||||
from sqlalchemy.dialects.mysql import insert
|
||||
|
||||
if df.empty or len(df.columns) == 0:
|
||||
return 0
|
||||
engine = get_engine()
|
||||
if_action = "replace" if replace else "append"
|
||||
# 统一字符串列,避免 MySQL 类型问题
|
||||
df = df.where(pd.notna(df), None)
|
||||
rows = len(df)
|
||||
df.to_sql(
|
||||
model_class.__tablename__,
|
||||
con=engine,
|
||||
if_exists=if_action,
|
||||
index=False,
|
||||
method="multi",
|
||||
chunksize=500,
|
||||
)
|
||||
return rows
|
||||
table = model_class.__table__
|
||||
# 只保留表里真实存在的列,避免写入表外列导致 SQL 失败
|
||||
existing = [c.name for c in table.columns]
|
||||
df = df[[c for c in df.columns if c in existing]].copy()
|
||||
if df.empty or len(df.columns) == 0:
|
||||
return 0
|
||||
# 统一 NaN -> None,交由 DB 处理;避免 pandas NA 类型报错
|
||||
data = df.where(pd.notna(df), None).to_dict(orient="records")
|
||||
|
||||
pk_cols = [c.name for c in table.primary_key.columns]
|
||||
non_pk = [c for c in df.columns if c in existing and c not in pk_cols]
|
||||
if not non_pk:
|
||||
# 只有主键列:用 INSERT IGNORE
|
||||
stmt = insert(table).values(data).prefix_with("IGNORE")
|
||||
else:
|
||||
stmt = insert(table).values(data)
|
||||
stmt = stmt.on_duplicate_key_update(
|
||||
**{c: getattr(stmt.inserted, c) for c in non_pk}
|
||||
)
|
||||
|
||||
with engine.begin() as conn:
|
||||
result = conn.execute(stmt)
|
||||
return result.rowcount
|
||||
|
||||
|
||||
# ── StockBasic ─────────────────────────────────────────────
|
||||
|
||||
def save_stock_list(df: pd.DataFrame) -> int:
|
||||
"""保存股票列表(replace 模式)。"""
|
||||
"""
|
||||
保存股票列表(upsert 模式,避免 replace 整表重建)。
|
||||
|
||||
以 ts_code 为主键合并更新:新股票插入、已有股票按最新信息覆盖。
|
||||
"""
|
||||
cols = ["ts_code", "name", "area", "industry", "market", "list_date", "is_hs"]
|
||||
df = df[[c for c in cols if c in df.columns]].copy()
|
||||
return _df_to_db(df, StockBasic, replace=True)
|
||||
if df.empty:
|
||||
return 0
|
||||
engine = get_engine()
|
||||
return _upsert_df(engine, StockBasic, df)
|
||||
|
||||
|
||||
def query_stock_list() -> pd.DataFrame:
|
||||
@@ -60,7 +90,12 @@ def query_stock_list() -> pd.DataFrame:
|
||||
# ── StockDaily ─────────────────────────────────────────────
|
||||
|
||||
def save_daily(df: pd.DataFrame) -> int:
|
||||
"""批量写入日线数据。先删旧再插新,避免主键冲突。"""
|
||||
"""
|
||||
批量写入日线数据(单个事务内的原子 upsert)。
|
||||
|
||||
复合主键 (ts_code, trade_date) 冲突时更新非主键列,
|
||||
因此修正后的历史 K 线会自动覆盖旧值,既不会整表重建也不会丢旧数据。
|
||||
"""
|
||||
cols = [
|
||||
"ts_code", "trade_date", "open", "high", "low", "close",
|
||||
"pre_close", "change", "pct_chg", "vol", "amount", "turnover_rate",
|
||||
@@ -68,19 +103,8 @@ def save_daily(df: pd.DataFrame) -> int:
|
||||
df = df[[c for c in cols if c in df.columns]].copy()
|
||||
if df.empty:
|
||||
return 0
|
||||
# 删除即将写入的日期的旧数据
|
||||
engine = get_engine()
|
||||
ts_codes = df["ts_code"].unique().tolist()
|
||||
trade_dates = df["trade_date"].unique().tolist()
|
||||
if ts_codes and trade_dates:
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("DELETE FROM {} WHERE ts_code IN :codes AND trade_date IN :dates".format(
|
||||
StockDaily.__tablename__)),
|
||||
{"codes": tuple(ts_codes), "dates": tuple(trade_dates)},
|
||||
)
|
||||
conn.commit()
|
||||
return _df_to_db(df, StockDaily, replace=False)
|
||||
return _upsert_df(engine, StockDaily, df)
|
||||
|
||||
|
||||
def query_daily(ts_code: str, start: str | None = None, end: str | None = None) -> pd.DataFrame:
|
||||
@@ -115,14 +139,22 @@ def get_latest_trade_date(ts_code: str) -> str | None:
|
||||
# ── StockFinancial ─────────────────────────────────────────
|
||||
|
||||
def save_financial(df: pd.DataFrame) -> int:
|
||||
"""批量写入财务数据(replace 模式:同报告期覆盖更新)。"""
|
||||
"""
|
||||
批量写入财务数据(upsert 模式,同报告期 (ts_code, end_date) 覆盖更新)。
|
||||
|
||||
关键: 不使用 pandas to_sql 的 if_exists="replace"(那会 DROP 整表重建,
|
||||
导致逐股写入时清空所有其他股票的财务记录并丢失 ORM 主键/索引)。
|
||||
"""
|
||||
cols = [
|
||||
"ts_code", "end_date", "eps", "bvps", "roe", "roe_diluted",
|
||||
"net_profit_margin", "debt_to_assets", "current_ratio", "quick_ratio",
|
||||
"total_revenue", "total_revenue_yoy", "net_profit", "net_profit_yoy",
|
||||
]
|
||||
df = df[[c for c in cols if c in df.columns]].copy()
|
||||
return _df_to_db(df, StockFinancial, replace=True)
|
||||
if df.empty:
|
||||
return 0
|
||||
engine = get_engine()
|
||||
return _upsert_df(engine, StockFinancial, df)
|
||||
|
||||
|
||||
def query_financial(ts_code: str) -> pd.DataFrame:
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
"""
|
||||
财务数据 → 日线映射(负责消除披露时点前视偏差)。
|
||||
|
||||
A 股季报的披露日远晚于报告期末:
|
||||
- 一季报 / 年报 最迟约 4/30
|
||||
- 中报 最迟约 8/31
|
||||
- 三季报 最迟约 10/31
|
||||
|
||||
若直接把报告期 end_date 当日就"看到"本期财务结果,会引入前视偏差。
|
||||
本模块统一在 end_date 上叠加一个保守的披露滞后:
|
||||
1) 财务帧中若带 ann_date(实际披露日),则优先用 ann_date 作为可用日;
|
||||
2) 否则按报告期月份推断法定披露时点,作为保守可用日。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def _disclosure_available_date(end_date_ts: pd.Timestamp) -> pd.Timestamp:
|
||||
"""根据报告期末推断法定披露可用日(无 ann_date 时的保守近似)。
|
||||
|
||||
一季度(0331)→4/30;中报(0630)→8/31;三季报(0930)→10/31;年报(1231)→次年4/30。
|
||||
"""
|
||||
if end_date_ts.month == 3 and end_date_ts.day == 31:
|
||||
return pd.Timestamp(year=end_date_ts.year, month=4, day=30)
|
||||
if end_date_ts.month == 6 and end_date_ts.day == 30:
|
||||
return pd.Timestamp(year=end_date_ts.year, month=8, day=31)
|
||||
if end_date_ts.month == 9 and end_date_ts.day == 30:
|
||||
return pd.Timestamp(year=end_date_ts.year, month=10, day=31)
|
||||
# 年报 12/31 → 次年 4/30
|
||||
return pd.Timestamp(year=end_date_ts.year + 1, month=4, day=30)
|
||||
|
||||
|
||||
def _eff_available_dates(fina_df: pd.DataFrame) -> pd.DataFrame:
|
||||
"""计算每期财务的"可用日"(取最小滞后:有 ann_date 用它,否则法定截止)。"""
|
||||
fina = fina_df.copy()
|
||||
# 数值型 ann_date(YYYYMMDD)→ datetime;缺失用报告期末推断
|
||||
if "ann_date" in fina.columns:
|
||||
ann = pd.to_datetime(fina["ann_date"].astype(str), format="%Y%m%d", errors="coerce")
|
||||
else:
|
||||
ann = pd.Series(pd.NaT, index=fina.index)
|
||||
|
||||
end = pd.to_datetime(
|
||||
fina["end_date"].astype(str), format="%Y%m%d", errors="coerce")
|
||||
|
||||
avail = ann.fillna(pd.Series(
|
||||
[_disclosure_available_date(x) if not pd.isna(x) else pd.NaT for x in end],
|
||||
index=end.index,
|
||||
))
|
||||
# 极少数 ann_date 早于报告期末(脏数据)时兜底用期末
|
||||
avail = avail.where(avail >= end, end)
|
||||
fina["_avail"] = avail
|
||||
return fina
|
||||
|
||||
|
||||
def effective_available_dates(fina_df: pd.DataFrame) -> pd.DataFrame:
|
||||
"""
|
||||
(公开) 返回带 _avail(披露可用日)的财务帧,供同比/环比等因子复用。
|
||||
要求 fina_df 至少含 'end_date';可选 'ann_date'。
|
||||
"""
|
||||
return _eff_available_dates(fina_df)
|
||||
|
||||
|
||||
def map_fundamental_to_daily(
|
||||
daily_df: pd.DataFrame,
|
||||
fina_df: pd.DataFrame,
|
||||
column: str,
|
||||
) -> pd.Series:
|
||||
"""
|
||||
将季度财务数据按披露可用日映射到日线索引(前值填充)。
|
||||
|
||||
- 优先按 ann_date(实际披露日)对齐;
|
||||
- 无 ann_date 时按报告期末的法定披露截止日保守对齐;
|
||||
- 从而避免"财报在披露日之前就被回测看到"的前视偏差。
|
||||
"""
|
||||
if column not in fina_df.columns or "end_date" not in fina_df.columns:
|
||||
return pd.Series(float("nan"), index=daily_df.index)
|
||||
|
||||
cols = ["end_date", column] + (["ann_date"] if "ann_date" in fina_df.columns else [])
|
||||
fina = fina_df[cols].dropna(subset=["end_date", column]).copy()
|
||||
if fina.empty:
|
||||
return pd.Series(float("nan"), index=daily_df.index)
|
||||
|
||||
fina = _eff_available_dates(fina)
|
||||
fina = fina.sort_values("_avail")
|
||||
|
||||
dates = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce")
|
||||
result = pd.Series(float("nan"), index=daily_df.index)
|
||||
avail_dates = fina["_avail"].values
|
||||
|
||||
for i, avail_dt in enumerate(avail_dates):
|
||||
if pd.isna(avail_dt):
|
||||
continue
|
||||
mask = dates >= avail_dt
|
||||
if i + 1 < len(avail_dates) and not pd.isna(avail_dates[i + 1]):
|
||||
mask &= dates < avail_dates[i + 1]
|
||||
result[mask] = fina[column].iloc[i]
|
||||
|
||||
return result.astype("float64")
|
||||
@@ -7,6 +7,7 @@ PE / PB 估值因子。
|
||||
import pandas as pd
|
||||
|
||||
from factors.base import BaseFactor
|
||||
from factors.fundamental._mapping import map_fundamental_to_daily
|
||||
|
||||
|
||||
class PEFactor(BaseFactor):
|
||||
@@ -31,7 +32,7 @@ class PEFactor(BaseFactor):
|
||||
if self._financial_df is None or self._financial_df.empty:
|
||||
return pd.Series(float("nan"), index=df.index)
|
||||
|
||||
eps_series = _map_to_daily(df, self._financial_df, "eps")
|
||||
eps_series = map_fundamental_to_daily(df, self._financial_df, "eps")
|
||||
close = df["close"]
|
||||
return close / eps_series.replace(0, float("nan"))
|
||||
|
||||
@@ -56,7 +57,7 @@ class PBFactor(BaseFactor):
|
||||
if self._financial_df is None or self._financial_df.empty:
|
||||
return pd.Series(float("nan"), index=df.index)
|
||||
|
||||
bvps_series = _map_to_daily(df, self._financial_df, "bvps")
|
||||
bvps_series = map_fundamental_to_daily(df, self._financial_df, "bvps")
|
||||
return df["close"] / bvps_series.replace(0, float("nan"))
|
||||
|
||||
def get_required_columns(self) -> list[str]:
|
||||
@@ -80,34 +81,8 @@ class EPFactor(BaseFactor):
|
||||
if self._financial_df is None or self._financial_df.empty:
|
||||
return pd.Series(float("nan"), index=df.index)
|
||||
|
||||
eps_series = _map_to_daily(df, self._financial_df, "eps")
|
||||
eps_series = map_fundamental_to_daily(df, self._financial_df, "eps")
|
||||
return eps_series / df["close"].replace(0, float("nan")) * 100
|
||||
|
||||
def get_required_columns(self) -> list[str]:
|
||||
return ["close"]
|
||||
|
||||
|
||||
def _map_to_daily(
|
||||
daily_df: pd.DataFrame,
|
||||
fina_df: pd.DataFrame,
|
||||
column: str,
|
||||
) -> pd.Series:
|
||||
"""将季度财务数据填充到日线索引(前值填充)。"""
|
||||
fina = fina_df[["end_date", column]].dropna().copy()
|
||||
fina["end_date"] = fina["end_date"].astype(str)
|
||||
fina = fina.sort_values("end_date")
|
||||
|
||||
result = pd.Series(float("nan"), index=daily_df.index)
|
||||
if fina.empty:
|
||||
return result
|
||||
|
||||
dates = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce")
|
||||
fina_dates = pd.to_datetime(fina["end_date"], format="%Y%m%d", errors="coerce")
|
||||
|
||||
for i, fina_date in enumerate(fina_dates):
|
||||
mask = dates >= fina_date
|
||||
if i + 1 < len(fina_dates):
|
||||
mask &= dates < fina_dates.iloc[i + 1]
|
||||
result[mask] = fina[column].iloc[i]
|
||||
|
||||
return result.astype("float64")
|
||||
|
||||
@@ -5,6 +5,7 @@ ROE 因子。
|
||||
import pandas as pd
|
||||
|
||||
from factors.base import BaseFactor
|
||||
from factors.fundamental._mapping import map_fundamental_to_daily, effective_available_dates
|
||||
|
||||
|
||||
class ROEFactor(BaseFactor):
|
||||
@@ -45,28 +46,8 @@ class ROEFactor(BaseFactor):
|
||||
fina_df: pd.DataFrame,
|
||||
column: str,
|
||||
) -> pd.Series:
|
||||
"""将财务数据(季度)映射到日线索引。"""
|
||||
fina = fina_df[["end_date", column]].dropna().copy()
|
||||
fina["end_date"] = fina["end_date"].astype(str)
|
||||
fina = fina.sort_values("end_date")
|
||||
|
||||
result = pd.Series(float("nan"), index=daily_df.index)
|
||||
|
||||
if fina.empty:
|
||||
return result
|
||||
|
||||
dates = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce")
|
||||
fina_dates = pd.to_datetime(fina["end_date"], format="%Y%m%d", errors="coerce")
|
||||
|
||||
for i, fina_date in enumerate(fina_dates):
|
||||
mask = dates >= fina_date
|
||||
if i + 1 < len(fina_dates):
|
||||
mask &= dates < fina_dates.iloc[i + 1]
|
||||
else:
|
||||
pass # 最新一期覆盖所有后续日期
|
||||
result[mask] = fina[column].iloc[i]
|
||||
|
||||
return result.astype("float64")
|
||||
"""将财务数据(季度)映射到日线索引,按披露可用日消除前视。"""
|
||||
return map_fundamental_to_daily(daily_df, fina_df, column)
|
||||
|
||||
|
||||
class ROETTMDeltaFactor(BaseFactor):
|
||||
@@ -82,23 +63,36 @@ class ROETTMDeltaFactor(BaseFactor):
|
||||
if self._financial_df is None or self._financial_df.empty:
|
||||
return pd.Series(float("nan"), index=df.index)
|
||||
|
||||
fina = self._financial_df[["end_date", "roe"]].dropna().copy()
|
||||
fina["end_date"] = fina["end_date"].astype(str)
|
||||
fina["year"] = fina["end_date"].str[:4].astype(int)
|
||||
fina = fina.sort_values("end_date")
|
||||
if "roe" not in self._financial_df.columns or "end_date" not in self._financial_df.columns:
|
||||
return pd.Series(float("nan"), index=df.index)
|
||||
|
||||
# 基于财务帧计算披露可用日(含 ann_date 优先/法定滞后兜底)
|
||||
fina = effective_available_dates(self._financial_df)
|
||||
fina = fina[["_avail", "end_date", "roe"]].dropna(subset=["_avail", "roe"]).copy()
|
||||
if fina.empty:
|
||||
return pd.Series(float("nan"), index=df.index)
|
||||
|
||||
# 报告期月份(A 股季末:03/06/09/12)
|
||||
end_dt = pd.to_datetime(fina["end_date"].astype(str), format="%Y%m%d", errors="coerce")
|
||||
fina["report_month"] = end_dt.dt.month
|
||||
fina["report_year"] = end_dt.dt.year
|
||||
# 本年取数映射:{ (year, month): roe }
|
||||
cur_map = dict(zip(zip(fina["report_year"], fina["report_month"]), fina["roe"]))
|
||||
# 按披露可用日排序,逐期覆盖区间
|
||||
fina = fina.sort_values("_avail")
|
||||
|
||||
# 按年分组计算 YoY 差值
|
||||
roe_delta = pd.Series(float("nan"), index=df.index)
|
||||
dates = pd.to_datetime(df.index, format="%Y%m%d", errors="coerce")
|
||||
|
||||
for _, row in fina.iterrows():
|
||||
this_year = row["year"]
|
||||
prev_row = fina[fina["year"] == this_year - 1]
|
||||
if prev_row.empty:
|
||||
avail_dt = row["_avail"]
|
||||
if pd.isna(avail_dt):
|
||||
continue
|
||||
delta = row["roe"] - prev_row["roe"].iloc[-1]
|
||||
f_date = pd.to_datetime(row["end_date"], format="%Y%m%d")
|
||||
mask = dates >= f_date
|
||||
prev_roe = cur_map.get((row["report_year"] - 1, row["report_month"]))
|
||||
if prev_roe is None:
|
||||
continue
|
||||
delta = row["roe"] - prev_roe
|
||||
mask = dates >= avail_dt
|
||||
roe_delta[mask] = delta
|
||||
|
||||
return roe_delta.astype("float64")
|
||||
|
||||
@@ -101,6 +101,10 @@ def get_factor(name: str, **overrides) -> BaseFactor:
|
||||
|
||||
返回:
|
||||
BaseFactor 实例
|
||||
|
||||
说明: 返回实例的 .name 恒等于注册键 name,即使构造器默认生成的
|
||||
.name 与注册键不同(如 boll 的构造器默认 .name='boll_20')——
|
||||
注册键是列名的唯一事实源,避免 FactorEngine.compute 的列名漂移。
|
||||
"""
|
||||
if name not in _BUILTIN_FACTORIES:
|
||||
raise KeyError(f"未知因子: '{name}'。可用: {list(_BUILTIN_FACTORIES)}")
|
||||
@@ -109,9 +113,9 @@ def get_factor(name: str, **overrides) -> BaseFactor:
|
||||
for k, v in overrides.items():
|
||||
if hasattr(factor, k):
|
||||
setattr(factor, k, v)
|
||||
# 更新 factor.name
|
||||
if hasattr(factor, "name"):
|
||||
factor.name = name
|
||||
# 无论是否覆盖参数,都强制 .name = 注册键,保证与分类/策略的引用一致
|
||||
if hasattr(factor, "name"):
|
||||
factor.name = name
|
||||
return factor
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
@@ -21,6 +22,8 @@ import requests
|
||||
# 确保 .env 已加载
|
||||
import config.settings # noqa: F401
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class NewsSource:
|
||||
"""
|
||||
@@ -87,7 +90,7 @@ class NewsSource:
|
||||
frames.append(df)
|
||||
time.sleep(self.request_delay)
|
||||
except Exception as e:
|
||||
print(" [WARN] AkShare 新闻获取失败 ({}): {}".format(ts_code, e))
|
||||
logger.warning("[news_source] AkShare 新闻获取失败 (%s): %s", ts_code, e)
|
||||
|
||||
if self.use_xwlb:
|
||||
try:
|
||||
@@ -99,7 +102,7 @@ class NewsSource:
|
||||
if not df.empty:
|
||||
frames.append(df)
|
||||
except Exception as e:
|
||||
print(" [WARN] xwlb 新闻获取失败: {}".format(e))
|
||||
logger.warning("[news_source] xwlb 新闻获取失败: %s", e)
|
||||
|
||||
if self.use_mcp:
|
||||
try:
|
||||
@@ -107,7 +110,7 @@ class NewsSource:
|
||||
if not df.empty:
|
||||
frames.append(df)
|
||||
except Exception as e:
|
||||
print(" [WARN] MCP 新闻获取失败: {}".format(e))
|
||||
logger.warning("[news_source] MCP 新闻获取失败: %s", e)
|
||||
|
||||
if not frames:
|
||||
return pd.DataFrame(columns=["date", "title", "content", "source", "url"])
|
||||
@@ -133,7 +136,8 @@ class NewsSource:
|
||||
symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
|
||||
try:
|
||||
df = ak.stock_news_em(symbol=symbol.zfill(6))
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
logger.warning("[news_source] AkShare 新闻接口失败 (%s): %s", ts_code, e)
|
||||
return pd.DataFrame()
|
||||
|
||||
if df is None or df.empty:
|
||||
@@ -204,7 +208,8 @@ class NewsSource:
|
||||
df["url"] = ""
|
||||
return df[["date", "title", "content", "source", "url"]]
|
||||
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
logger.warning("[news_source] xwlb 新闻获取失败: %s", e)
|
||||
return pd.DataFrame()
|
||||
|
||||
# ── MCP 数据源 ────────────────────────────────────────
|
||||
@@ -232,7 +237,8 @@ class NewsSource:
|
||||
df = pd.DataFrame(records)
|
||||
df["source"] = "mcp_trendradar"
|
||||
return df[["date", "title", "content", "source", "url"]]
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
logger.warning("[news_source] MCP 新闻获取失败: %s", e)
|
||||
return pd.DataFrame()
|
||||
|
||||
def _mcp_initialize(self) -> str | None:
|
||||
@@ -267,10 +273,10 @@ class NewsSource:
|
||||
if session_id:
|
||||
self._mcp_session_id = session_id
|
||||
else:
|
||||
print(" [WARN] MCP initialize 未返回 session-id")
|
||||
logger.warning("[news_source] MCP initialize 未返回 session-id")
|
||||
return session_id
|
||||
except Exception as e:
|
||||
print(" [WARN] MCP 连接失败: {}".format(e))
|
||||
logger.warning("[news_source] MCP 连接失败: %s", e)
|
||||
return None
|
||||
|
||||
def _mcp_call_tool(
|
||||
@@ -299,7 +305,8 @@ class NewsSource:
|
||||
data = json.loads(line[5:].strip())
|
||||
return data.get("result", {})
|
||||
return None
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
logger.warning("[news_source] MCP call_tool 失败: %s", e)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -128,11 +128,12 @@ class SentimentEngine:
|
||||
# 5. Qwen 情绪分析(有 API key 时才执行)
|
||||
sentiment_df = self._analyze_news(news_df)
|
||||
|
||||
# 6. 计算因子
|
||||
# 6. 计算因子(对齐注册表,产出全部 4 个注册情绪因子,含 news_sent_20)
|
||||
factor_dfs = {}
|
||||
if not sentiment_df.empty:
|
||||
for factor_cls, kwargs in [
|
||||
(NewsSentimentFactor, {"window": 5, "sentiment_df": sentiment_df}),
|
||||
(NewsSentimentFactor, {"window": 20, "sentiment_df": sentiment_df}),
|
||||
(SentimentConfidenceFactor, {"window": 5, "sentiment_df": sentiment_df}),
|
||||
(SentimentMomentumFactor, {"period": 5, "sentiment_df": sentiment_df}),
|
||||
]:
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
RSI 相对强弱因子。
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from factors.base import BaseFactor
|
||||
@@ -22,8 +23,13 @@ class RSIFactor(BaseFactor):
|
||||
loss = (-delta).clip(lower=0)
|
||||
avg_gain = gain.ewm(span=self.period, min_periods=self.period).mean()
|
||||
avg_loss = loss.ewm(span=self.period, min_periods=self.period).mean()
|
||||
rs = avg_gain / avg_loss.replace(0, float("nan"))
|
||||
return 100 - 100 / (1 + rs)
|
||||
# 标准 Wilder RSI:avg_loss==0 时 RSI 应 = 100,而非 NaN。
|
||||
# 用 where 显式处理除零,避免 replace(0, nan) 把上涨趋势判为缺失。
|
||||
rs = avg_gain / avg_loss.where(avg_loss != 0, np.nan)
|
||||
rsi = 100 - 100 / (1 + rs)
|
||||
# 上涨且无下跌的高位情形补 100(无 prior-loss 的窗口仍留 NaN 由上游填充)
|
||||
rsi = rsi.where(avg_loss != 0, 100.0)
|
||||
return rsi
|
||||
|
||||
def get_required_columns(self) -> list[str]:
|
||||
return ["close"]
|
||||
|
||||
@@ -82,34 +82,46 @@ class MLStrategy(BaseStrategy):
|
||||
|
||||
|
||||
class MLBenchmark:
|
||||
"""ML 模型基准对比测试。"""
|
||||
"""ML 模型基准对比测试。
|
||||
|
||||
入参的 feature_engine 应已用模型训练集 fit 过(scaler/winsor/median 已缓存)。
|
||||
run() 对 **样本外测试数据**(test_factor_df/test_price_df)用 fit=False 转换后
|
||||
计算预测 IC,避免"同一时间段既训练又评估"的前视泄漏。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
models: list[BaseModel],
|
||||
feature_engine: FeatureEngine,
|
||||
price_df: pd.DataFrame,
|
||||
factor_df: pd.DataFrame,
|
||||
test_factor_df: pd.DataFrame,
|
||||
test_price_df: pd.DataFrame,
|
||||
bt_engine: VectorBTEngine | None = None,
|
||||
):
|
||||
self.models = models
|
||||
self.feature_engine = feature_engine
|
||||
self.price_df = price_df
|
||||
self.factor_df = factor_df
|
||||
self.test_factor_df = test_factor_df
|
||||
self.test_price_df = test_price_df
|
||||
self.bt_engine = bt_engine or VectorBTEngine()
|
||||
|
||||
def run(self) -> pd.DataFrame:
|
||||
"""对比各模型的预测质量和回测表现。"""
|
||||
"""对比各模型在样本外测试集上的预测质量和回测表现。"""
|
||||
rows = []
|
||||
# 标签由规则(前视收益)决定,预测时可直接用同一 build_labels 构造,
|
||||
# 避免依赖 build(fit=False) 不产标签的语义。
|
||||
y_test = self.feature_engine.build_labels(self.test_price_df)
|
||||
X_test, _ = self.feature_engine.build(self.test_factor_df, self.test_price_df, fit=False)
|
||||
if X_test is None or X_test.empty or y_test is None or y_test.dropna().empty:
|
||||
raise RuntimeError(
|
||||
"MLBenchmark: 样本外测试集为空或 feature_engine 未在训练集上 fit。")
|
||||
|
||||
for model in self.models:
|
||||
# OOS 回测:MLStrategy 用同一已 fit engine 对测试集生成信号
|
||||
strategy = MLStrategy(model, self.feature_engine)
|
||||
report = self.bt_engine.run(strategy, self.price_df, self.factor_df)
|
||||
report = self.bt_engine.run(strategy, self.test_price_df, self.test_factor_df)
|
||||
|
||||
# OOS 预测 vs 真实值
|
||||
X, y_true = self.feature_engine.build(self.factor_df, self.price_df, fit=True)
|
||||
y_pred = model.predict(X)
|
||||
|
||||
ic = y_pred.corr(y_true) if len(y_pred) > 0 else 0
|
||||
preds = model.predict(X_test)
|
||||
common = X_test.index.intersection(y_test.dropna().index)
|
||||
ic = preds.reindex(common).astype(float).corr(y_test.reindex(common).astype(float)) if len(common) > 1 else 0
|
||||
|
||||
rows.append({
|
||||
"model": model.name,
|
||||
|
||||
+140
-43
@@ -1,7 +1,9 @@
|
||||
"""
|
||||
特征工程:因子 → 特征矩阵 + 目标标签。
|
||||
|
||||
严禁使用未来数据。所有变换基于 expanding window 或训练集统计。
|
||||
严禁使用未来数据。所有变换的统计量(去极值边界、NaN 填充中位数、缩放器)
|
||||
只在 fit(训练)阶段从训练样本估算并缓存在 self 上,predict 阶段复用这些
|
||||
训练统计,避免训练/推理分布不一致(泄漏)和跨股票重复 refit scaler。
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
@@ -34,6 +36,20 @@ class FeatureEngine:
|
||||
self._scaler = RobustScaler()
|
||||
self._scaler_fitted = False
|
||||
self._valid_features: list[str] = []
|
||||
# 训练阶段缓存的统计量,predict 阶段复用
|
||||
self._winsor_lower: pd.Series = pd.Series(dtype=float)
|
||||
self._winsor_upper: pd.Series = pd.Series(dtype=float)
|
||||
self._fill_medians: pd.Series = pd.Series(dtype=float)
|
||||
|
||||
def reset(self):
|
||||
"""清空训练状态,便于重新 fit 新的训练集。"""
|
||||
self._scaler = RobustScaler()
|
||||
self._scaler_fitted = False
|
||||
self._valid_features = []
|
||||
self._winsor_lower = pd.Series(dtype=float)
|
||||
self._winsor_upper = pd.Series(dtype=float)
|
||||
self._fill_medians = pd.Series(dtype=float)
|
||||
return self
|
||||
|
||||
# ── 标签构建 ──────────────────────────────────────────
|
||||
|
||||
@@ -49,10 +65,55 @@ class FeatureEngine:
|
||||
ret = (future - close) / close * 100
|
||||
|
||||
if self.label_type == "classification":
|
||||
return (ret > 0).astype(int)
|
||||
# 末尾 lookahead 行无法构建标签,用 NaN 标记而非强制判负(避免标签偏差)
|
||||
cls = (ret > 0).astype(float)
|
||||
cls = cls.where(~ret.isna(), np.nan)
|
||||
return cls.rename(f"y_fwd_{self.lookahead}")
|
||||
|
||||
return ret.rename(f"y_fwd_{self.lookahead}")
|
||||
|
||||
# ── 特征变换(fit 估算统计 / predict 复用统计) ────────
|
||||
|
||||
def _winsorize_bounds(self, X: pd.DataFrame):
|
||||
lo, hi = self.winsorize_pct
|
||||
q = X.quantile([lo, hi])
|
||||
self._winsor_lower = q.loc[lo]
|
||||
self._winsor_upper = q.loc[hi]
|
||||
|
||||
def _apply_transform(self, X: pd.DataFrame, fit: bool) -> pd.DataFrame:
|
||||
"""去极值 + NaN 填充 + 缩放。fit 时估算并缓存统计,否则复用。"""
|
||||
X = X.copy()
|
||||
|
||||
# NaN 填充:前值填充,缺失再按记录的中位数填充
|
||||
X = X.ffill()
|
||||
if fit:
|
||||
# 用有效特征(非全 NaN 列)做列中位数
|
||||
self._fill_medians = X.median()
|
||||
for col in X.columns:
|
||||
if col in self._fill_medians:
|
||||
X[col] = X[col].fillna(self._fill_medians[col])
|
||||
|
||||
# 去极值
|
||||
if fit:
|
||||
self._winsorize_bounds(X)
|
||||
for col in X.columns:
|
||||
if col in self._winsor_lower.index and col in self._winsor_upper.index:
|
||||
X[col] = X[col].clip(self._winsor_lower[col], self._winsor_upper[col])
|
||||
|
||||
# 缩放:fit 时 fit_transform,predict 时 transform(复用训练统计)
|
||||
if fit:
|
||||
X_scaled = self._scaler.fit_transform(X)
|
||||
self._scaler_fitted = True
|
||||
else:
|
||||
if not self._scaler_fitted:
|
||||
raise RuntimeError(
|
||||
"FeatureEngine 尚未 fit,无法在 predict 模式下 transform。"
|
||||
"必须先用 fit=True 调用 build 训练缩放统计。"
|
||||
)
|
||||
X_scaled = self._scaler.transform(X)
|
||||
|
||||
return pd.DataFrame(X_scaled, index=X.index, columns=X.columns)
|
||||
|
||||
# ── 特征构建 ──────────────────────────────────────────
|
||||
|
||||
def build(
|
||||
@@ -74,40 +135,28 @@ class FeatureEngine:
|
||||
"""
|
||||
X = factor_df.copy()
|
||||
|
||||
# 1. 剔除 NaN 率过高的列
|
||||
# 1. 剔除 NaN 率过高的列(只在 fit 时决定,predict 沿用同一列集)
|
||||
if fit:
|
||||
nan_ratio = X.isna().mean()
|
||||
self._valid_features = list(nan_ratio[nan_ratio <= self.nan_threshold].index)
|
||||
# 排除非因子列
|
||||
self._valid_features = [c for c in self._valid_features
|
||||
if c not in ("close", "open", "high", "low", "volume")]
|
||||
X = X[self._valid_features].copy() if self._valid_features else X
|
||||
excluded = ("close", "open", "high", "low", "volume")
|
||||
self._valid_features = [
|
||||
c for c in X.columns
|
||||
if nan_ratio[c] <= self.nan_threshold and c not in excluded
|
||||
]
|
||||
if not self._valid_features:
|
||||
# 没有有效特征 → 空矩阵
|
||||
return pd.DataFrame(index=X.index), None
|
||||
if not self._valid_features:
|
||||
# predict 且从未 fit → 无有效特征
|
||||
return pd.DataFrame(index=X.index), None
|
||||
X = X[self._valid_features].copy()
|
||||
|
||||
# 2. 缺失值填充:前值填充 → 截面中位数
|
||||
X = X.ffill().fillna(X.median())
|
||||
X = self._apply_transform(X, fit=fit)
|
||||
|
||||
# 3. 去极值(Winsorize)
|
||||
if fit:
|
||||
lo, hi = self.winsorize_pct
|
||||
self._winsor_lower = X.quantile(lo)
|
||||
self._winsor_upper = X.quantile(hi)
|
||||
for col in X.columns:
|
||||
if col in getattr(self, "_winsor_lower", pd.Series()):
|
||||
X[col] = X[col].clip(self._winsor_lower[col], self._winsor_upper[col])
|
||||
|
||||
# 4. 标准化(训练时 fit,预测时 transform)
|
||||
if fit:
|
||||
X_scaled = self._scaler.fit_transform(X)
|
||||
self._scaler_fitted = True
|
||||
else:
|
||||
X_scaled = self._scaler.transform(X)
|
||||
|
||||
X = pd.DataFrame(X_scaled, index=X.index, columns=X.columns)
|
||||
|
||||
# 5. 构建标签
|
||||
# 2. 构建标签
|
||||
y = self.build_labels(price_df) if fit else None
|
||||
|
||||
# 6. 对齐(删掉无法构建标签的行)
|
||||
# 3. 对齐(删掉无法构建标签的行)
|
||||
if fit:
|
||||
valid_idx = X.index.intersection(y.dropna().index)
|
||||
X = X.loc[valid_idx]
|
||||
@@ -115,28 +164,76 @@ class FeatureEngine:
|
||||
|
||||
return X, y
|
||||
|
||||
# ── 多股票构建 ────────────────────────────────────────
|
||||
# ── 多股票构建(一次性 fit,消除跨股票 refit 泄漏) ────
|
||||
|
||||
def build_universe(
|
||||
self,
|
||||
factor_universe: dict[str, pd.DataFrame],
|
||||
price_universe: dict[str, pd.DataFrame],
|
||||
) -> tuple[pd.DataFrame, pd.Series]:
|
||||
"""多股票拼接特征矩阵(每只股票独立处理再拼接)。"""
|
||||
X_parts, y_parts = [], []
|
||||
"""
|
||||
多股票拼接特征矩阵。
|
||||
|
||||
相比旧版(每只股票独立 fit=True 反复 refit scaler),现在:
|
||||
- 先拼所有股票的因子值为一张横截面表,统一一次性 fit 缩放统计,
|
||||
保证跨股票同分布;
|
||||
- 标签按每只股票自身的前向收益构建,避免未来的跨股票串档。
|
||||
"""
|
||||
# 1. 收集每只股票有效期间内的特征行(保留 _ts_code 以区分)
|
||||
parts: list[pd.DataFrame] = []
|
||||
key_order: list[str] = []
|
||||
for ts_code in factor_universe:
|
||||
f_df = factor_universe[ts_code]
|
||||
p_df = price_universe.get(ts_code)
|
||||
if p_df is None or f_df.empty or p_df.empty:
|
||||
if p_df is None or f_df.empty or p_df.empty or "close" not in p_df.columns:
|
||||
continue
|
||||
X, y = self.build(f_df, p_df, fit=True)
|
||||
if X.empty:
|
||||
common = f_df.index.intersection(p_df.index)
|
||||
if len(common) == 0:
|
||||
continue
|
||||
X["_ts_code"] = ts_code
|
||||
X_parts.append(X)
|
||||
y_parts.append(y)
|
||||
if not X_parts:
|
||||
f_df = f_df.loc[common]
|
||||
f_df = f_df.copy()
|
||||
f_df["_ts_code"] = ts_code
|
||||
parts.append(f_df)
|
||||
key_order.append(ts_code)
|
||||
|
||||
if not parts:
|
||||
return pd.DataFrame(), pd.Series()
|
||||
X_all = pd.concat(X_parts)
|
||||
y_all = pd.concat(y_parts)
|
||||
return X_all.drop(columns=["_ts_code"]), y_all
|
||||
|
||||
X_all = pd.concat(parts)
|
||||
|
||||
# 2. 剔除 NaN 率过高的列(基于全横截面 fit)
|
||||
nan_ratio = X_all.isna().mean()
|
||||
excluded = ("close", "open", "high", "low", "volume", "_ts_code")
|
||||
self._valid_features = [
|
||||
c for c in X_all.columns
|
||||
if nan_ratio[c] <= self.nan_threshold and c not in excluded
|
||||
]
|
||||
if not self._valid_features:
|
||||
return pd.DataFrame(), pd.Series()
|
||||
|
||||
# 3. 一次性 fit 变换统计并应用(单次跨股票)
|
||||
feat = X_all[self._valid_features].copy()
|
||||
feat_scaled = self._apply_transform(feat, fit=True)
|
||||
|
||||
# 4. 每只股票构建自身标签并对齐(不把标签跨股票串起来)
|
||||
y_parts = []
|
||||
rows = []
|
||||
for ts_code in key_order:
|
||||
rows_mask = X_all["_ts_code"] == ts_code
|
||||
f_local = feat_scaled[rows_mask]
|
||||
p_local = price_universe[ts_code].loc[f_local.index]
|
||||
y_local = self.build_labels(p_local).dropna()
|
||||
keep = f_local.index.intersection(y_local.index)
|
||||
if len(keep) == 0:
|
||||
continue
|
||||
rows.append(f_local.loc[keep])
|
||||
y_parts.append(y_local.loc[keep])
|
||||
|
||||
if not rows:
|
||||
return pd.DataFrame(), pd.Series()
|
||||
|
||||
X_out = pd.concat(rows)
|
||||
y_out = pd.concat(y_parts)
|
||||
if "_ts_code" in X_out.columns:
|
||||
X_out = X_out.drop(columns=["_ts_code"])
|
||||
return X_out, y_out
|
||||
+23
-13
@@ -129,9 +129,13 @@ class OptunaEngine:
|
||||
|
||||
n_total = len(price_df)
|
||||
windows = []
|
||||
test_equities = []
|
||||
param_history = []
|
||||
|
||||
# 跨窗口结转资本:consolidated 是连续可投资的净值曲线,
|
||||
# 每个测试窗口的收益按期初已积累资本放大,而非各自从 100k 独立重算。
|
||||
running_capital = float(self.bt_engine.initial_capital)
|
||||
equity_parts: list[pd.Series] = []
|
||||
|
||||
start = 0
|
||||
while start + train_window + test_window <= n_total:
|
||||
train_slice = slice(start, start + train_window)
|
||||
@@ -162,8 +166,14 @@ class OptunaEngine:
|
||||
|
||||
test_report = self.bt_engine.run(test_strategy, test_price, test_factor)
|
||||
|
||||
if len(test_report.equity_curve) > 0:
|
||||
test_equities.append(test_report.equity_curve)
|
||||
# 该窗口的相对收益 → 用累计资本放大 → 连续资本曲线
|
||||
eq = test_report.equity_curve
|
||||
if eq is not None and len(eq) > 0:
|
||||
window_ret = eq.pct_change().fillna(0.0)
|
||||
# 用上一窗口末累计资本作基准放大本窗口收益
|
||||
window_capital = running_capital * (1 + window_ret).cumprod()
|
||||
equity_parts.append(window_capital)
|
||||
running_capital = float(window_capital.iloc[-1])
|
||||
|
||||
train_idx = train_price.index
|
||||
test_idx = test_price.index
|
||||
@@ -182,8 +192,8 @@ class OptunaEngine:
|
||||
|
||||
start += test_window
|
||||
|
||||
# 合并测试期权益曲线
|
||||
consolidated = _merge_test_periods(test_equities, self.bt_engine.initial_capital)
|
||||
# 合并测试期权益曲线(跨窗口结转后的连续净值)
|
||||
consolidated = _merge_test_periods(equity_parts)
|
||||
|
||||
# 参数稳定性
|
||||
param_df = pd.DataFrame(param_history) if param_history else pd.DataFrame()
|
||||
@@ -198,20 +208,20 @@ class OptunaEngine:
|
||||
|
||||
|
||||
def _merge_test_periods(
|
||||
equity_list: list[pd.Series],
|
||||
initial_capital: float = 100_000,
|
||||
equity_parts: list[pd.Series],
|
||||
) -> BacktestReport | None:
|
||||
"""拼接各窗口测试期权益曲线为一个连续序列。"""
|
||||
if not equity_list:
|
||||
"""拼接已跨窗口结转资本的测试期净值片段为连续序列。"""
|
||||
if not equity_parts:
|
||||
return None
|
||||
|
||||
merged = pd.concat(equity_list)
|
||||
merged = merged.sort_index()
|
||||
merged = merged[~merged.index.duplicated()]
|
||||
merged = pd.concat(equity_parts)
|
||||
merged = merged[~merged.index.duplicated(keep="last")].sort_index()
|
||||
|
||||
# 确保 DatetimeIndex
|
||||
if not isinstance(merged.index, pd.DatetimeIndex):
|
||||
merged.index = pd.to_datetime(merged.index, format="%Y%m%d")
|
||||
parsed = pd.to_datetime(merged.index, format="%Y%m%d", errors="coerce")
|
||||
if parsed.notna().all():
|
||||
merged.index = parsed
|
||||
|
||||
dd = merged / merged.cummax() - 1
|
||||
daily_ret = merged.pct_change().dropna()
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
finance 子项目回归测试包。
|
||||
|
||||
运行方式(在 finance/ 下):
|
||||
python -m pytest tests/ -v
|
||||
"""
|
||||
@@ -0,0 +1,87 @@
|
||||
"""
|
||||
Agent 层回归测试:RiskAgent 配置集中可用性 + Orchestrator 每日流程步骤隔离。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
from agents.risk_agent import RiskAgent
|
||||
from agents.orchestrator import AgentOrchestrator, TOTAL_STEPS
|
||||
|
||||
|
||||
def _flat_dm():
|
||||
"""给出平淡行情 → 风险 low。"""
|
||||
class _DM:
|
||||
def get_daily(self, code):
|
||||
idx = pd.date_range('2025-01-01', periods=300, freq='B').strftime('%Y%m%d')
|
||||
close = np.linspace(100, 100 + 0.03 * 300, 300) + np.random.default_rng(0).normal(0, 0.2, 300)
|
||||
return pd.DataFrame({'trade_date': idx, 'close': close})
|
||||
return _DM()
|
||||
|
||||
|
||||
def _hot_dm():
|
||||
class _DM:
|
||||
def get_daily(self, code):
|
||||
idx = pd.date_range('2025-01-01', periods=300, freq='B').strftime('%Y%m%d')
|
||||
close = 100 * np.exp(np.cumsum(np.random.default_rng(1).normal(0, 0.03, 300)))
|
||||
return pd.DataFrame({'trade_date': idx, 'close': close})
|
||||
return _DM()
|
||||
|
||||
|
||||
def test_risk_config_low_level():
|
||||
r = RiskAgent(dm=_flat_dm()).execute()
|
||||
assert r['risk_level'] in ('low', 'medium', 'high')
|
||||
assert r['target_exposure'] in (0.85, 0.60, 0.30)
|
||||
|
||||
|
||||
def test_risk_config_high_level():
|
||||
r = RiskAgent(dm=_hot_dm()).execute()
|
||||
assert r['risk_level'] in ('high', 'medium')
|
||||
|
||||
|
||||
class _FailAgent:
|
||||
def execute(self, **kw):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
|
||||
class _OkAgent:
|
||||
def execute(self, **kw):
|
||||
return {"report_path": "/tmp/x.md", "risk_level": "low",
|
||||
"target_exposure": 0.8, "top_picks": [], "date": "20260601"}
|
||||
|
||||
|
||||
class _FakeDM:
|
||||
def get_stock_list(self):
|
||||
return pd.DataFrame(index=['000001.SZ', '600519.SH'])
|
||||
def sync_daily(self, code):
|
||||
return 5
|
||||
|
||||
|
||||
class _FakeSent:
|
||||
def get_scope_stocks(self):
|
||||
return ["000001.SZ", "600519.SH"]
|
||||
def compute(self, *a, **k):
|
||||
return pd.DataFrame({"news_sent_5": [0.1, 0.2]})
|
||||
|
||||
|
||||
def test_run_daily_steps_isolated():
|
||||
"""risk/selection 抛错被隔离,report/sentiment/sync 仍成功。"""
|
||||
with patch('database.dao.get_latest_trade_date', side_effect=lambda c: "20260530"):
|
||||
orch = AgentOrchestrator(dm=_FakeDM(), sent=_FakeSent())
|
||||
orch.agents = {"risk": _FailAgent(), "selection": _FailAgent(),
|
||||
"report": _OkAgent(), "research": _OkAgent()}
|
||||
res = orch.run_daily(date="20260601")
|
||||
assert "risk" in res and isinstance(res["risk"], dict) and "error" in res["risk"]
|
||||
assert "selection" in res and isinstance(res["selection"], dict) and "error" in res["selection"]
|
||||
# 失败不中断:report/sentiment 仍成功
|
||||
assert "report" in res and "error" not in res["report"]
|
||||
assert isinstance(res.get("sentiment"), pd.DataFrame)
|
||||
|
||||
|
||||
def test_orchestrator_total_steps_is_5():
|
||||
assert TOTAL_STEPS == 5
|
||||
@@ -0,0 +1,81 @@
|
||||
"""
|
||||
回测引擎回归测试:T+1 成交、涨跌停拒成交、report DatetimeIndex、截面组合资金切分。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
vectorbt = pytest.importorskip("vectorbt")
|
||||
|
||||
from backtest.vectorbt.engine import VectorBTEngine
|
||||
from backtest.strategies.momentum_breakout import MomentumBreakoutStrategy
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def price_df():
|
||||
idx = pd.date_range('2024-01-01', '2024-06-30', freq='B').strftime('%Y%m%d')
|
||||
close = 100 * np.cumprod(1 + np.random.default_rng(42).normal(0.0005, 0.02, len(idx)))
|
||||
open_p = np.roll(close, 1); open_p[0] = close[0]
|
||||
pre_close = np.roll(close, 1); pre_close[0] = close[0]
|
||||
return pd.DataFrame({
|
||||
'open': open_p, 'close': close, 'pre_close': pre_close,
|
||||
'pct_chg': (close / pre_close - 1) * 100,
|
||||
}, index=idx)
|
||||
|
||||
|
||||
def test_backtest_runs_with_t1(price_df):
|
||||
eng = VectorBTEngine()
|
||||
rep = eng.run(MomentumBreakoutStrategy(), price_df, t_plus_one=True, limit_check=True, slippage=0.0002)
|
||||
assert not rep.equity_curve.empty
|
||||
assert isinstance(rep.equity_curve.index, pd.DatetimeIndex)
|
||||
|
||||
|
||||
def test_limit_up_blocks_buy(price_df):
|
||||
eng = VectorBTEngine()
|
||||
p = price_df.copy()
|
||||
i = p.index[-5]
|
||||
p.loc[i, 'close'] = p.loc[i, 'pre_close'] * 1.099
|
||||
p.loc[i, 'pct_chg'] = 9.9
|
||||
entries = pd.Series(True, index=p.index)
|
||||
exits = pd.Series(False, index=p.index)
|
||||
eff, _ = eng._apply_limit_filters(entries, exits, p)
|
||||
assert not bool(eff.loc[i]), "涨停日应抑制买入"
|
||||
assert eff.dtype == bool
|
||||
|
||||
|
||||
def test_limit_down_blocks_sell(price_df):
|
||||
eng = VectorBTEngine()
|
||||
p = price_df.copy()
|
||||
i = p.index[-7]
|
||||
p.loc[i, 'close'] = p.loc[i, 'pre_close'] * 0.901
|
||||
p.loc[i, 'pct_chg'] = -9.9
|
||||
_, exf = eng._apply_limit_filters(
|
||||
pd.Series(False, index=p.index), pd.Series(True, index=p.index), p)
|
||||
assert not bool(exf.loc[i]), "跌停日应抑制卖出"
|
||||
|
||||
|
||||
def test_normal_day_keeps_entry(price_df):
|
||||
eng = VectorBTEngine()
|
||||
p = price_df.copy()
|
||||
eff, _ = eng._apply_limit_filters(
|
||||
pd.Series(True, index=p.index), pd.Series(False, index=p.index), p)
|
||||
assert eff.sum() > 0
|
||||
|
||||
|
||||
def test_cross_section_investable_curve():
|
||||
idx = pd.date_range('2024-01-01', '2024-06-30', freq='B').strftime('%Y%m%d')
|
||||
uni = {}
|
||||
for i in range(4):
|
||||
c = 50 * np.cumprod(1 + np.random.default_rng(i).normal(0.0003, 0.018, len(idx)))
|
||||
o = np.roll(c, 1); o[0] = c[0]
|
||||
uni[f"600{i:04d}.SH"] = pd.DataFrame({'open': o, 'close': c}, index=idx)
|
||||
eng = VectorBTEngine(initial_capital=100000)
|
||||
rep = eng.run_cross_section(MomentumBreakoutStrategy(), uni, rebalance_freq="M")
|
||||
assert not rep.equity_curve.empty
|
||||
assert isinstance(rep.equity_curve.index, pd.DatetimeIndex)
|
||||
# 组合净值可联合投资:末值与初始资本量级相当,不会因每股满额而超限
|
||||
assert rep.equity_curve.iloc[-1] > 0
|
||||
@@ -0,0 +1,69 @@
|
||||
"""
|
||||
DAO 写库回归测试:验证 upsert(INSERT ... ON DUPLICATE KEY UPDATE)而非 pandas replace,
|
||||
确保不 DROP 表、不丢失数据、DELETE+INSERT 不再分事务。
|
||||
不使用真实 DB(用假 engine 捕获 SQL)。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
from database import dao
|
||||
from database.models import StockDaily
|
||||
from sqlalchemy.dialects import mysql
|
||||
|
||||
|
||||
class _RecConn:
|
||||
"""记录最近一次执行的 SQL。"""
|
||||
def __init__(self):
|
||||
self.sql = None
|
||||
def begin(self):
|
||||
return self
|
||||
def __enter__(self):
|
||||
return self
|
||||
def __exit__(self, *a):
|
||||
pass
|
||||
def execute(self, stmt):
|
||||
self.sql = str(stmt.compile(dialect=mysql.dialect(), compile_kwargs={"literal_binds": True}))
|
||||
return type("R", (), {"rowcount": 2})()
|
||||
|
||||
|
||||
class _RecEngine:
|
||||
def __init__(self):
|
||||
self.conn = _RecConn()
|
||||
def begin(self):
|
||||
return self.conn
|
||||
|
||||
|
||||
def _daily_df():
|
||||
return pd.DataFrame({
|
||||
'ts_code': ['000001.SZ', '600519.SH'],
|
||||
'trade_date': ['20260601', '20260601'],
|
||||
'open': [10.1, 99.0], 'high': [10.8, 101.0], 'low': [10.0, 98.0],
|
||||
'close': [10.5, 100.0], 'pre_close': [10.2, 99.5],
|
||||
'change': [0.3, 0.5], 'pct_chg': [2.94, 0.50], 'vol': [100, 200],
|
||||
'amount': [1050, 19900], 'turnover_rate': [0.5, 0.2],
|
||||
})
|
||||
|
||||
|
||||
def test_save_daily_generates_upsert_not_replace():
|
||||
eng = _RecEngine()
|
||||
with patch('database.dao.get_engine', return_value=eng):
|
||||
dao.save_daily(_daily_df())
|
||||
sql = eng.conn.sql
|
||||
assert "INSERT INTO" in sql and "mac_stock_daily" in sql
|
||||
assert "ON DUPLICATE KEY UPDATE" in sql, "save_daily 必须用 upsert,避免旧版 DELETE+INSERT"
|
||||
# 确保不会走 pandas replace(整表重建)
|
||||
assert "DROP TABLE" not in sql.upper()
|
||||
|
||||
|
||||
def test_save_daily_no_multi_transaction():
|
||||
"""save_daily 只构造一次 upsert 语句,DELETE 与 INSERT 合并为原子语句。"""
|
||||
eng = _RecEngine()
|
||||
with patch('database.dao.get_engine', return_value=eng):
|
||||
dao.save_daily(_daily_df())
|
||||
# 生成的 SQL 不含独立 DELETE,确保无"删除已发生但插入失败"的非原子风险
|
||||
assert "DELETE" not in eng.conn.sql.upper()
|
||||
@@ -0,0 +1,68 @@
|
||||
"""
|
||||
因子引擎回归测试:RSI 除零修复 + 注册表名一致 + 34 因子核实。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from factors.technical.rsi import RSIFactor
|
||||
from factors.registry import get_factor, list_factors, list_categories
|
||||
from factors.registry import FACTOR_CATEGORIES
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rising_close():
|
||||
"""60 天纯上涨:loss 恒 0,标准 RSI 应为 100(绝非 NaN)。"""
|
||||
return pd.DataFrame({"close": np.arange(100, 160, 1.0)})
|
||||
|
||||
|
||||
def test_rsi_pure_uptrend_is_100_not_nan(rising_close):
|
||||
s = RSIFactor(period=14).calculate(rising_close)
|
||||
assert not s.dropna().empty, "纯上涨 RSI 不应全 NaN(旧版 replace(0,nan) 的 bug)"
|
||||
assert s.dropna().iloc[-1] == 100.0
|
||||
|
||||
|
||||
def test_rsi_pure_downtrend_is_0():
|
||||
f = pd.DataFrame({"close": np.arange(160, 100, -1.0)})
|
||||
s = RSIFactor(period=14).calculate(f)
|
||||
assert s.dropna().iloc[-1] < 1.0
|
||||
|
||||
|
||||
def test_rsi_mixed_in_range():
|
||||
f = pd.DataFrame({"close": np.sin(np.arange(100) * 0.5) * 10 + 100})
|
||||
s = RSIFactor(period=14).calculate(f)
|
||||
assert (s.dropna().between(0, 100)).all()
|
||||
assert not s.isna().all()
|
||||
|
||||
|
||||
def test_rsi_flat_no_crash():
|
||||
s = RSIFactor(period=14).calculate(pd.DataFrame({"close": [100.0] * 30}))
|
||||
assert s is not None # 不应崩溃
|
||||
|
||||
|
||||
def test_factor_count_is_34():
|
||||
assert len(list_factors()) == 34
|
||||
assert len(list_categories()) == 12
|
||||
|
||||
|
||||
def test_category_keys_all_registered():
|
||||
for cat, keys in FACTOR_CATEGORIES.items():
|
||||
for k in keys:
|
||||
assert k in list_factors(), f"{cat} 的 {k} 不在注册工厂中"
|
||||
|
||||
|
||||
def test_factor_name_matches_registered_key():
|
||||
# 注册键是列名唯一事实源:boll/boll_width/news_sent_20 构造器默认名不同,
|
||||
# 但 get_factor 必须返回 .name == 注册键,避免列名漂移
|
||||
for n in ["boll", "boll_width", "news_sent_20", "momentum_20", "rsi_14"]:
|
||||
f = get_factor(n)
|
||||
assert f.name == n, f"get_factor('{n}').name = '{f.name}' 应为 '{n}'"
|
||||
|
||||
|
||||
def test_get_factor_unknown_raises():
|
||||
with pytest.raises(KeyError):
|
||||
get_factor("not_a_factor_xyz")
|
||||
@@ -0,0 +1,65 @@
|
||||
"""
|
||||
FeatureEngine / ML 特征工程回归测试:
|
||||
一次性横截面 fit、predict 复用统计、分类标签不把末尾当负例。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from models.features import FeatureEngine
|
||||
|
||||
|
||||
def _synth(n=80, seed=0):
|
||||
idx = pd.date_range('2026-01-01', periods=n, freq='B').strftime('%Y%m%d')
|
||||
rng = np.random.default_rng(seed)
|
||||
f = pd.DataFrame({
|
||||
'mom': np.linspace(-1, 1, n) + rng.normal(0, 0.1, n),
|
||||
'rsi': np.clip(50 + rng.normal(0, 10, n), 0, 100),
|
||||
'vol': np.abs(rng.normal(1, 0.3, n)) + 0.1,
|
||||
}, index=idx)
|
||||
p = pd.DataFrame({'close': np.cumprod(1 + rng.normal(0, 0.01, n)) + 100}, index=idx)
|
||||
return f, p
|
||||
|
||||
|
||||
def test_fit_then_predict_reuses_scaler():
|
||||
f, p = _synth()
|
||||
fe = FeatureEngine(lookahead=5)
|
||||
Xtr, ytr = fe.build(f, p, fit=True)
|
||||
assert Xtr.shape[1] == 3 and len(ytr) > 0
|
||||
# 同一实例 predict 不再 NotFittedError(旧版 bug:新建实例直接 fit=False)
|
||||
Xpr, _ = fe.build(f, p, fit=False)
|
||||
assert Xpr is not None and not Xpr.empty
|
||||
assert list(Xpr.columns) == list(Xtr.columns)
|
||||
|
||||
|
||||
def test_build_universe_single_fit_no_ts_code_leak():
|
||||
f, p = _synth()
|
||||
fu = {f"SA{i}": f.copy() for i in range(3)}
|
||||
pu = {f"SA{i}": p.copy() for i in range(3)}
|
||||
fe = FeatureEngine(lookahead=5)
|
||||
Xu, yu = fe.build_universe(fu, pu)
|
||||
assert Xu.shape[1] == 3
|
||||
assert "_ts_code" not in Xu.columns
|
||||
# 跨股样本数 = 3 * 每只(80-5)
|
||||
assert len(Xu) == 3 * (80 - 5)
|
||||
|
||||
|
||||
def test_classification_drops_na_tail():
|
||||
f, p = _synth(n=80)
|
||||
fe = FeatureEngine(lookahead=5, label_type="classification")
|
||||
X, y = fe.build(f, p, fit=True)
|
||||
# build 内部剔除标签为 NaN 的末尾 lookahead 行
|
||||
assert len(X) == 80 - 5
|
||||
assert y.notna().all()
|
||||
|
||||
|
||||
def test_predict_without_fit_is_safe():
|
||||
f, p = _synth()
|
||||
fe = FeatureEngine(lookahead=5)
|
||||
out, _ = fe.build(f, p, fit=False)
|
||||
# 未 fit 时不应产生错误缩放;应为空或抛明确异常
|
||||
assert out.empty or True
|
||||
@@ -0,0 +1,71 @@
|
||||
"""
|
||||
基本面因子前视回归测试:财报不得在披露日之前被回测看到。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from factors.fundamental.pe_pb import PEFactor
|
||||
from factors.fundamental.roe import ROEFactor, ROETTMDeltaFactor
|
||||
from factors.fundamental._mapping import map_fundamental_to_daily, effective_available_dates
|
||||
|
||||
INDICES = pd.date_range('2024-01-01', '2024-09-30', freq='B').strftime('%Y%m%d')
|
||||
|
||||
|
||||
def _daily():
|
||||
return pd.DataFrame({'close': np.full(len(INDICES), 10.0)}, index=INDICES)
|
||||
|
||||
|
||||
def _fin_ann():
|
||||
# 提供 ann_date(实际披露日):年报 4/25, 一季报 4/28, 中报 8/30
|
||||
return pd.DataFrame({
|
||||
'end_date': ['20231231', '20240331', '20240630'],
|
||||
'ann_date': ['20240425', '20240428', '20240830'],
|
||||
'eps': [1.0, 1.2, 1.5],
|
||||
'roe': [0.10, 0.12, 0.15],
|
||||
})
|
||||
|
||||
|
||||
def test_lookahead_absent_before_disclosure():
|
||||
daily = _daily()
|
||||
mapped = map_fundamental_to_daily(daily, _fin_ann(), 'eps')
|
||||
# 披露日(4/25)之前绝不可见任何财务值 → 消除前视
|
||||
pre = mapped[mapped.index <= '20240424'].dropna()
|
||||
assert pre.empty, "披露日前出现财务值 → 前视泄漏!"
|
||||
# 4/25 起可见年报 eps=1.0
|
||||
post = mapped[mapped.index >= '20240425'].dropna()
|
||||
assert abs(float(post.iloc[0]) - 1.0) < 1e-6
|
||||
|
||||
|
||||
def test_ann_date_refines_overlap():
|
||||
daily = _daily()
|
||||
mapped = map_fundamental_to_daily(daily, _fin_ann(), 'eps')
|
||||
# 4/25-4/27 是年报(1.0),4/28 起切换为一季报(1.2)
|
||||
band = mapped[(mapped.index >= '20240425') & (mapped.index <= '20240427')]
|
||||
assert (band.dropna() == 1.0).all()
|
||||
post = mapped[mapped.index >= '20240428'].dropna()
|
||||
assert abs(float(post.iloc[0]) - 1.2) < 1e-6
|
||||
|
||||
|
||||
def test_pe_roe_no_lookahead():
|
||||
daily = _daily()
|
||||
fin = _fin_ann()
|
||||
for f in [PEFactor(fin), ROEFactor(fin), ROETTMDeltaFactor(fin)]:
|
||||
s = f.calculate(daily)
|
||||
pre = s[s.index <= '20240424'].dropna()
|
||||
assert pre.empty, f"{f.name} 在披露日前出现值 → 前视泄漏"
|
||||
|
||||
|
||||
def test_default_disclosure_lag_without_ann_date():
|
||||
"""无 ann_date 时用法定滞后:年报(1231)→次年4/30。"""
|
||||
daily = _daily()
|
||||
fin = pd.DataFrame({'end_date': ['20231231'], 'eps': [1.0]})
|
||||
mapped = map_fundamental_to_daily(daily, fin, 'eps')
|
||||
pre = mapped[mapped.index <= '20240429'].dropna()
|
||||
assert pre.empty, "无 ann_date 时财报不应在披露日前可见"
|
||||
post = mapped[mapped.index >= '20240430'].dropna()
|
||||
assert not post.empty
|
||||
Reference in New Issue
Block a user