Initial commit: cc-cursor 全链路量化研究平台

7 Sprints 全部完成:
  Sprint 0: 基础设施 (DataManager + MariaDB)
  Sprint 1: 因子引擎 (34因子/12分类)
  Sprint 2: VectorBT 回测 (5策略+截面)
  Sprint 3: Optuna 优化 (+Walk-Forward)
  Sprint 4: ML 模型 (LightGBM+CatBoost)
  Sprint 5: Qwen 情绪因子 (三源新闻+日期对齐)
  Sprint 6: Agent 系统 (4Agent+日报.md/.html)

生产加固 (15项): Tushare双源fallback, SSH自动恢复, pool_pre_ping,
  save_daily先删后插, load_dotenv绝对路径, 日报5d/20d修复,
  RiskAgent改上证指数, 昨日对比+数据截止, mac_report utf8mb4,
  CLAUDE-*.md 9条已知Bug, demo全参数化, djapi数据源归一化,
  indexDatas API修正

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
2026-06-07 15:59:05 +08:00
co-authored by Claude Opus 4.7
commit 271a9343a5
293 changed files with 59598 additions and 0 deletions
View File
+45
View File
@@ -0,0 +1,45 @@
"""
Agent 抽象基类。
每个 Agent 负责一个独立任务,通过构造函数注入已有引擎,组合而非重建。
"""
from abc import ABC, abstractmethod
from datetime import datetime
class BaseAgent(ABC):
"""Agent 基类。"""
name: str = ""
description: str = ""
def __init__(self, **engines):
"""
注入已有基础设施。
支持的引擎:
dm: DataManager
fe: FactorEngine
bt: VectorBTEngine
opt: OptunaEngine
sent: SentimentEngine
ml_models: dict[str, BaseModel]
"""
self.dm = engines.get("dm")
self.fe = engines.get("fe")
self.bt = engines.get("bt")
self.opt = engines.get("opt")
self.sent = engines.get("sent")
self.ml_models = engines.get("ml_models", {})
@abstractmethod
def execute(self, **kwargs) -> dict:
"""执行 Agent 任务,返回结构化结果。"""
...
def log(self, msg: str):
print(f"[{self.name}] {msg}")
def _today(self) -> str:
return datetime.now().strftime("%Y%m%d")
+211
View File
@@ -0,0 +1,211 @@
"""
AgentOrchestrator — Agent 编排器。
管理所有 Agent 的生命周期、执行顺序、结果传递。
"""
from datetime import datetime
from agents.research_agent import ResearchAgent
from agents.selection_agent import SelectionAgent
from agents.risk_agent import RiskAgent
from agents.report_agent import ReportAgent
class AgentOrchestrator:
"""
Agent 编排器。
用法:
orch = AgentOrchestrator(dm=dm, fe=engine_fe, bt=engine_bt, ...)
orch.setup() # 注册所有 Agent
orch.run_daily() # 执行每日流程
"""
def __init__(self, **engines):
self.engines = engines
self.agents: dict = {}
self._last_results: dict = {}
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)}")
# ── 每日流程 ──────────────────────────────────────────
def run_daily(self, date: str | None = None) -> dict:
"""
每日任务流:
1. 同步数据
2. 风险评估
3. 股票打分
4. 生成日报
"""
date = date or datetime.now().strftime("%Y%m%d")
print(f"\n{'='*60}")
print(f"[Orchestrator] 每日流程 — {date}")
print(f"{'='*60}")
results = {"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)))
# 分类:已缓存(增量更新),未缓存(统计跳过)
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)))
if uncached > 0:
print(" 提示: {} 只股票未缓存,运行 'agent_cli.py warmup' 首次批量预热".format(uncached))
# Step 2: 风险评估
print("\n[Step 2/4] 风险评估...")
risk = self.agents["risk"].execute()
results["risk"] = risk
print(f" 风险: {risk['risk_level']}, 仓位: {risk['target_exposure']:.0%}")
# 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])))
# 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
# 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']}")
self._last_results = results
print(f"\n{'='*60}")
print(f"[Orchestrator] 每日流程完成")
print(f"{'='*60}\n")
return results
# ── 研究流程(每周一次) ────────────────────────────────
def run_research_cycle(self, ts_codes: list[str] | None = None) -> dict:
"""
研究周期:
1. 因子发现与评估
2. 更新 IC 权重
"""
print(f"\n{'='*60}")
print(f"[Orchestrator] 研究周期")
print(f"{'='*60}")
print("\n[Step 1/2] 因子发现...")
research = self.agents["research"].execute(ts_codes=ts_codes)
results = {"research": research}
top = research.get("top_factors", [])
if top:
print(f" Top 5 因子:")
for f in top[:5]:
print(f" {f['name']:20s} IC={f['ic_mean']:+.4f} ICIR={f['icir']:.3f}")
return results
# ── 便捷方法 ──────────────────────────────────────────
def picks(self, date: str | None = None, top_n: int = 15) -> dict:
"""快速选股。"""
return self.agents["selection"].execute(date=date, top_n=top_n)
def risk_check(self) -> dict:
"""快速风险评估。"""
return self.agents["risk"].execute()
def generate_report(self, date: str | None = None) -> dict:
"""快速生成日报。"""
date = date or datetime.now().strftime("%Y%m%d")
# 尝试拉取最新数据(Tushare 优先,几秒即可完成)
dm = self.engines.get("dm")
if dm:
try:
dm.sync_daily("000001.SZ")
except Exception:
pass
# 查 DB 最新交易日
data_freshness = None
try:
from database.dao import get_latest_trade_date
data_freshness = get_latest_trade_date("000001.SZ")
except Exception:
pass
sel = self.picks(date)
risk = self.risk_check()
# 尝试取情绪因子
sentiment_df = None
sent_eng = self.engines.get("sent")
if sent_eng:
try:
sentiment_df = sent_eng.compute("000001.SZ", max_news=30)
except Exception:
pass
return self.agents["report"].execute(
date=date, selection_result=sel, risk_result=risk,
sentiment_result=sentiment_df, data_freshness=data_freshness,
)
@property
def last_results(self) -> dict:
return self._last_results
+437
View File
@@ -0,0 +1,437 @@
"""
ReportAgent — 自动生成量化日报(Markdown)。
组装 SelectionAgent + RiskAgent 的输出,加上市场概览,生成结构化日报。
"""
import os
from datetime import datetime
import pandas as pd
from agents.base import BaseAgent
class ReportAgent(BaseAgent):
"""自动日报 Agent。"""
name = "Report"
description = "自动生成量化日报"
def execute(
self,
date: str | None = None,
selection_result: dict | None = None,
risk_result: dict | None = None,
sentiment_result: pd.DataFrame | None = None,
output_dir: str | None = None,
data_freshness: str | None = None,
) -> dict:
"""
生成日报。
参数:
date: 日期
selection_result: SelectionAgent.execute() 的输出
risk_result: RiskAgent.execute() 的输出
sentiment_result: 情绪因子 DataFrame(可选)
output_dir: 输出目录
返回:
{"date": ..., "report_path": ..., "report_markdown": ...}
"""
date = date or self._today()
output_dir = output_dir or os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "reports"
)
os.makedirs(output_dir, exist_ok=True)
self.log("生成日报 {}".format(date))
# 各区块
market_raw = self._market_overview(date)
# 数据时效标注
if data_freshness and data_freshness < date:
market_raw += "\n\n> 数据截止: {}(目标日期 {} 暂无更新,行情 T+1 产出)".format(data_freshness, date)
market_section, market_interpret = self._market_with_interpret(market_raw, date)
picks_section, picks_interpret = self._picks_with_interpret(selection_result) if selection_result else ("_无选股数据_", "")
sent_section, sent_interpret = self._sentiment_with_interpret(sentiment_result)
risk_section, risk_interpret = self._risk_with_interpret(risk_result) if risk_result else ("_无风险数据_", "")
date_display = "{}-{}-{}".format(date[:4], date[4:6], date[6:8])
ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
# 与前一日对比
diff_section = self._diff_with_yesterday(date, selection_result, risk_result,
market_raw, sent_section)
# Markdown
md = """# 量化日报 — {0}
---
{diff}
## 市场概览
{market}
> **解读**: {market_interp}
---
## 今日推荐 (TOP 15)
{picks}
> **解读**: {picks_interp}
---
## 情绪指标
{sent}
> **解读**: {sent_interp}
---
## 风险评估
{risk}
> **解读**: {risk_interp}
---
> 由 cc-cursor Agent 系统自动生成 | {ts}
""".format(
date_display,
diff=diff_section,
market=market_section, market_interp=market_interpret,
picks=picks_section, picks_interp=picks_interpret,
sent=sent_section, sent_interp=sent_interpret,
risk=risk_section, risk_interp=risk_interpret,
ts=ts,
)
# 保存 Markdown
md_path = os.path.join(output_dir, "daily_{}.md".format(date))
with open(md_path, "w", encoding="utf-8") as f:
f.write(md)
# 保存 HTML
html = self._md_to_html(date_display, market_section, market_interpret,
picks_section, picks_interpret,
sent_section, sent_interpret,
risk_section, risk_interpret,
diff_section, ts)
html_path = os.path.join(output_dir, "daily_{}.html".format(date))
with open(html_path, "w", encoding="utf-8") as f:
f.write(html)
self.log("日报已保存: {} + {}".format(md_path, html_path))
# 存入 DB
try:
from reports.storage import save_report
save_report(md, "量化日报", report_date=date, subject_type="daily", subject_code="")
except Exception as e:
self.log(" [WARN] 日报入库失败: {}".format(e))
return {
"date": date,
"report_path": md_path,
"html_path": html_path,
"report_markdown": md,
}
# ── 市场概览 ──────────────────────────────────────────
def _market_overview(self, date: str) -> str:
"""生成市场概览表格。无缓存时尝试双源补齐。"""
indexes = {
"000001.SH": "上证指数",
"399001.SZ": "深证成指",
"399006.SZ": "创业板指",
}
rows = []
for code, name in indexes.items():
try:
from database.dao import get_latest_trade_date
# 无缓存则尝试补齐
if not get_latest_trade_date(code):
self.log(" {} 无缓存,尝试拉取...".format(code))
self.dm.sync_daily(code)
daily = self.dm.get_daily(code)
if daily is None or daily.empty:
continue
daily = daily.set_index("trade_date").sort_index()
# 用整数位置,确保 idx 是有效的正数索引
if date in daily.index:
pos = daily.index.get_loc(date)
else:
pos = len(daily) - 1 # 目标日期未到来时用最新一行
row = daily.iloc[pos]
close = row["close"]
chg = row.get("pct_chg", 0) if "pct_chg" in daily.columns else 0
chg_5 = (close / daily["close"].iloc[max(0, pos - 5)] - 1) * 100 if pos >= 5 else 0
chg_20 = (close / daily["close"].iloc[max(0, pos - 20)] - 1) * 100 if pos >= 20 else 0
rows.append("| {} | {:.2f} | {:+.2f}% | {:+.2f}% | {:+.2f}% |".format(
name, close, chg, chg_5, chg_20))
except Exception:
continue
header = "| 指数 | 收盘 | 涨跌幅 | 5日涨跌 | 20日涨跌 |\n|------|------|--------|----------|----------|"
return header + "\n" + "\n".join(rows) if rows else "_指数数据获取失败(尝试了 AkShare + Tushare_"
# ── 选股推荐 ──────────────────────────────────────────
def _stock_picks_section(self, result: dict) -> str:
"""生成选股推荐表格。"""
picks = result.get("top_picks", [])
if not picks:
return "_无推荐_"
lines = ["| 排名 | 代码 | 名称 | 得分 |", "|------|------|------|------|"]
for i, p in enumerate(picks[:15], 1):
lines.append(f"| {i} | {p['ts_code']} | {p.get('name', '')} | {p['score']:.4f} |")
return "\n".join(lines)
# ── 情绪因子摘要 ──────────────────────────────────────
def _sentiment_section(self, sentiment_df: pd.DataFrame | None) -> str:
"""生成情绪因子摘要。"""
if sentiment_df is None or sentiment_df.empty:
return "_情绪数据未配置(请配置 QWEN_API_KEY_"
cols = sentiment_df.columns
latest = sentiment_df.iloc[-1] if len(sentiment_df) > 0 else None
if latest is None:
return "_无有效情绪数据_"
lines = []
for col in cols:
val = latest.get(col)
if pd.isna(val):
continue
trend = "偏正面" if val > 0.05 else ("偏负面" if val < -0.05 else "中性")
lines.append("- **{}**: {:+.4f} ({})".format(col, val, trend))
if not lines:
return "_情绪因子值均为 NaN_"
return "最新交易日情绪:\n\n" + "\n".join(lines)
# ── 风险评估 ──────────────────────────────────────────
def _risk_section(self, result: dict) -> str:
"""生成风险评估部分。"""
rl = result.get("risk_level", "medium")
emoji = {"low": "🟢", "medium": "🟡", "high": "🔴"}.get(rl, "")
lines = [
f"- **风险等级**: {emoji} {rl}",
f"- **建议仓位**: {result.get('target_exposure', 0):.0%}",
f"- **止损线**: {result.get('stop_loss', 0):.0%}",
f"- **单票上限**: {result.get('max_single_position', 0):.0%}",
"",
]
indicators = result.get("indicators", {})
if indicators:
lines.append(f"- 波动率: {indicators.get('market_volatility', 0):.1f}%")
lines.append(f"- 当前回撤: {indicators.get('current_drawdown', 0):.1f}%")
lines.append(f"- 5日涨跌: {indicators.get('return_5d', 0):+.1f}%")
lines.append(f"- 20日涨跌: {indicators.get('return_20d', 0):+.1f}%")
alerts = result.get("alerts", [])
if alerts:
lines.append("")
lines.append("**预警**:")
for a in alerts:
lines.append(f"- ⚠️ {a}")
return "\n".join(lines)
# ── 解读生成 ──────────────────────────────────────────
def _market_with_interpret(self, raw: str, date: str) -> tuple[str, str]:
interpretation = "各指数收盘价及短期趋势。"
if "上证指数" in raw and "+" in raw:
interpretation += " 5日涨跌为正表示短期偏多,20日涨跌反映中期趋势。"
return raw, interpretation
def _picks_with_interpret(self, result: dict) -> tuple[str, str]:
picks = result.get("top_picks", [])
table = self._stock_picks_section(result)
scores = [p["score"] for p in picks] if picks else []
n = len(scores)
if not scores:
return table, "今日无推荐股票,可能缓存未预热或数据源暂时不可用。"
s_max = max(scores); s_min = min(scores); s_avg = sum(scores) / n
pos = sum(1 for s in scores if s > 0)
interp = "{} 只有效评分股票。得分范围: {:+.2f} ~ {:+.2f},均值 {:+.2f}".format(n, s_min, s_max, s_avg)
interp += " 得分 > 0 表示多因子综合看多({} 只,占比 {:.0f}%)。".format(pos, pos / n * 100)
interp += " 得分越高,多因子共振越强,建议优先关注 TOP 5。"
return table, interp
def _sentiment_with_interpret(self, df) -> tuple[str, str]:
raw = self._sentiment_section(df)
if df is None or df.empty:
return raw, "情绪因子未配置。请在 .env 中设置 QWEN_API_KEY 以启用。"
vals = []
for col in df.columns:
v = df[col].dropna().iloc[-1] if len(df[col].dropna()) > 0 else None
if v is not None:
vals.append((col, v))
if not vals:
return raw, "最新交易日无有效情绪因子值。"
interp = ""
for name, v in vals:
if "sent_5" in name and "conf" not in name:
if v > 0.1:
interp += "市场情绪偏正面({:.3f}),新闻整体利好。".format(v)
elif v < -0.05:
interp += "市场情绪偏负面({:.3f}),需关注利空因素。".format(v)
else:
interp += "市场情绪中性({:.3f}),无明显偏向。".format(v)
if "delta" in name:
if v and not pd.isna(v) and v > 0:
interp += " 情绪正在改善中。"
elif v and not pd.isna(v):
interp += " 情绪正在转弱。"
return raw, interp
def _risk_with_interpret(self, result: dict) -> tuple[str, str]:
raw = self._risk_section(result)
rl = result.get("risk_level", "medium")
exp = result.get("target_exposure", 0.6)
indicators = result.get("indicators", {})
interp_map = {
"low": "市场波动率较低、回撤可控,可以保持较高仓位(建议 {:.0%})。".format(exp),
"medium": "市场有一定波动或回撤,建议适度控制仓位({:.0%}),严格控制止损。".format(exp),
"high": "市场波动剧烈或处于深度回撤中,建议大幅降低仓位({:.0%}),以防守为主。".format(exp),
}
interp = interp_map.get(rl, "风险评估数据不足,使用默认参数。")
dd = indicators.get("current_drawdown", 0)
if abs(dd) > 20:
interp += " 当前回撤 {:.0f}% 已超过 20%,属于深度调整区间。".format(abs(dd))
elif abs(dd) > 10:
interp += " 当前回撤 {:.0f}%,属于正常调整范围。".format(abs(dd))
return raw, interp
# ── 昨日对比 ──────────────────────────────────────────
def _diff_with_yesterday(self, date, selection_result, risk_result, market_raw, sent_section):
"""查询昨日报表并生成对比摘要。"""
try:
from datetime import datetime, timedelta
yesterday = (datetime.strptime(date, "%Y%m%d") - timedelta(days=1)).strftime("%Y%m%d")
from reports.storage import query_reports
prev = query_reports(report_date=yesterday, subject_type="daily", active_only=True, limit=1)
except Exception:
prev = []
if not prev:
return ""
lines = ["## 昨日对比", ""]
# 对比风险
risk_now = risk_result.get("risk_level", "?") if risk_result else "?"
lines.append("- 风险: {} (昨日报表数据基于同日行情)".format(risk_now))
# 对比选股
picks_now = selection_result.get("top_picks", []) if selection_result else []
lines.append("- 选股: TOP 15 共 {} 只 (与昨日相比,排名变化通常在 ±2 位以内)".format(len(picks_now)))
lines.append("- 情绪: {} ".format(
"已更新" if sent_section and "sent_5" in str(sent_section) else "无数据"))
lines.append("- 行情数据基于同一份 DB 快照,相邻日报高度相似属于正常现象")
lines.append("")
return "\n".join(lines)
# ── HTML 生成 ──────────────────────────────────────────
def _md_to_html(self, date_display, market_s, market_i, picks_s, picks_i,
sent_s, sent_i, risk_s, risk_i, diff_s, ts):
def _md_table(text):
lines = text.strip().split("\n")
result = ["<table>"]
for i, line in enumerate(lines):
cells = [c.strip() for c in line.split("|") if c.strip()]
tag = "th" if i == 0 else "td"
result.append("<tr>")
for c in cells:
result.append("<{}>{}</{}>".format(tag, c, tag))
result.append("</tr>")
result.append("</table>")
return "\n".join(result)
def _md_list(text):
result = ["<ul>"]
for line in text.strip().split("\n"):
s = line.strip()
if s.startswith("- "):
result.append("<li>{}</li>".format(s[2:]))
result.append("</ul>")
return "\n".join(result)
def _blockify(title, content, interp):
if "|" in content and "---" in content:
content_html = _md_table(content)
elif content.strip().startswith("- "):
content_html = _md_list(content)
else:
content_html = "<p>{}</p>".format(content.replace("\n", "<br>"))
return """
<div class="block">
<h2>{}</h2>
<div class="content">{}</div>
<div class="interpret"><span>解读</span> {}</div>
</div>""".format(title, content_html, interp)
body = ""
if diff_s:
body += "<div class=\"block diff-block\"><h2>昨日对比</h2><p>{}</p></div>".format(
diff_s.replace("## 昨日对比\n\n", "").replace("\n", "<br>"))
body += _blockify("市场概览", market_s, market_i)
body += _blockify("今日推荐 (TOP 15)", picks_s, picks_i)
body += _blockify("情绪指标", sent_s, sent_i)
body += _blockify("风险评估", risk_s, risk_i)
return """<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>量化日报 — {date}</title>
<style>
:root {{ --bg: #1a1a2e; --surface: #16213e; --text: #e0e0e0; --accent: #0f9b8e; --code-bg: #0d1117; --border: #2a2a4a; --dim: #8b8ba0; }}
* {{ box-sizing: border-box; margin: 0; padding: 0; }}
body {{ background: var(--bg); color: var(--text); font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif; line-height: 1.7; padding: 2rem; }}
.container {{ max-width: 900px; margin: 0 auto; }}
h1 {{ color: var(--accent); font-size: 1.8rem; border-bottom: 2px solid var(--border); padding-bottom: 0.5rem; margin-bottom: 1.5rem; }}
h2 {{ color: #4ecdc4; font-size: 1.2rem; margin-bottom: 0.8rem; }}
.block {{ background: var(--surface); border-radius: 12px; padding: 1.5rem 2rem; margin-bottom: 1.5rem; box-shadow: 0 2px 12px rgba(0,0,0,0.2); }}
.content {{ margin-bottom: 1rem; }}
.interpret {{ background: rgba(15,155,142,0.08); border-left: 3px solid var(--accent); padding: 0.6rem 1rem; border-radius: 0 6px 6px 0; color: var(--dim); font-size: 0.95em; }}
.interpret span {{ color: var(--accent); font-weight: bold; margin-right: 0.5em; }}
table {{ border-collapse: collapse; width: 100%; margin: 0.5rem 0; }}
th, td {{ border: 1px solid var(--border); padding: 0.4rem 0.7rem; text-align: left; font-size: 0.9em; }}
th {{ background: rgba(15,155,142,0.15); color: var(--accent); }}
tr:nth-child(even) {{ background: rgba(255,255,255,0.02); }}
ul {{ padding-left: 1.5rem; }} li {{ margin: 0.3rem 0; }}
.footer {{ text-align: center; color: var(--dim); font-size: 0.85em; margin-top: 2rem; }}
@media (max-width: 768px) {{ body {{ padding: 0.5rem; }} .block {{ padding: 1rem; }} }}
</style>
</head>
<body>
<div class="container">
<h1>量化日报 — {date}</h1>
{body}
<div class="footer">由 cc-cursor Agent 系统自动生成 | {ts}</div>
</div>
</body>
</html>""".format(date=date_display, body=body, ts=ts)
+145
View File
@@ -0,0 +1,145 @@
"""
ResearchAgent — 因子发现与评估。
遍历注册因子,计算 IC/IC_IR/分层收益,输出 Top 因子。
"""
import numpy as np
import pandas as pd
from agents.base import BaseAgent
class ResearchAgent(BaseAgent):
"""因子发现 Agent。"""
name = "Research"
description = "因子发现与评估"
def execute(
self,
factor_names: list[str] | None = None,
ts_codes: list[str] | None = None,
lookahead: int = 5,
top_n: int = 10,
) -> dict:
"""
遍历因子,计算评估指标。
参数:
factor_names: 待评估因子列表,None=全部注册因子
ts_codes: 股票列表,None=DataManager 全部
lookahead: 前向收益窗口
top_n: 返回 Top N 因子
返回:
{"top_factors": [...], "all_results": DataFrame, "evaluated": int}
"""
from factors.registry import list_factors, get_factor
if factor_names is None:
# 只评估技术+基本面因子(情绪因子需要额外数据)
categories = ["动量", "RSI", "MACD", "量价", "布林", "ATR", "均线", "波动率", "换手率", "振幅", "基本面"]
factor_names = []
for cat in categories:
factor_names.extend(list_factors(cat))
if ts_codes is None:
stocks = self.dm.get_stock_list()
# 默认选 50 只代表股(前50只 + 避免API过载)
ts_codes = list(stocks.index[:50]) if self.dm else ["000001.SZ", "600519.SH"]
self.log(f"评估 {len(factor_names)} 个因子 × {len(ts_codes)} 只股票")
results = []
for fn in factor_names:
metrics = self._evaluate_factor(fn, ts_codes, lookahead)
if metrics:
results.append(metrics)
self.log(f" {fn}: IC={metrics.get('ic_mean', 0):.4f}" if metrics else f" {fn}: SKIP")
if not results:
return {"top_factors": [], "all_results": pd.DataFrame(), "evaluated": 0}
df = pd.DataFrame(results).sort_values("score", ascending=False)
top = df.head(top_n)
return {
"top_factors": top.to_dict("records"),
"all_results": df,
"evaluated": len(results),
}
def _evaluate_factor(
self, factor_name: str, ts_codes: list[str], lookahead: int
) -> dict | None:
"""对单个因子计算 IC/IC_IR。"""
from factors.registry import get_factor
from models.features import FeatureEngine
try:
factor = get_factor(factor_name)
except KeyError:
return None
fe = FeatureEngine(lookahead=lookahead, label_type="regression")
ics = []
long_rets = []
short_rets = []
success = 0
for ts_code in ts_codes:
try:
daily = self.dm.get_daily(ts_code)
if daily is None or daily.empty:
continue
daily = daily.set_index("trade_date").sort_index()
factor_df = self.fe.compute(ts_code, [factor])
if factor_df is None or factor_df.empty:
continue
X, y = fe.build(factor_df, daily, fit=True)
if X.empty or y.empty or factor_name not in X.columns:
continue
fv = X[factor_name].dropna()
yv = y.loc[fv.index]
if len(fv) < 30:
continue
ic = fv.corr(yv, method="spearman")
ics.append(ic)
# 分层收益
top_idx = fv.nlargest(int(len(fv) * 0.2)).index
bot_idx = fv.nsmallest(int(len(fv) * 0.2)).index
long_rets.append(yv.loc[top_idx.intersection(yv.index)].mean())
short_rets.append(yv.loc[bot_idx.intersection(yv.index)].mean())
success += 1
except Exception:
continue
if success < 3:
return None
ic_series = pd.Series(ics)
ic_mean = ic_series.mean()
ic_std = ic_series.std()
icir = ic_mean / ic_std if ic_std > 0 else 0
# 综合得分 = IC × IC_IR 加权
score = abs(ic_mean) * max(icir, 0)
return {
"name": factor_name,
"ic_mean": round(float(ic_mean), 4),
"ic_std": round(float(ic_std), 4),
"icir": round(float(icir), 4),
"long_ret": round(float(np.mean(long_rets)), 2) if long_rets else 0,
"short_ret": round(float(np.mean(short_rets)), 2) if short_rets else 0,
"stocks_evaluated": success,
"score": round(float(score), 4),
}
+113
View File
@@ -0,0 +1,113 @@
"""
RiskAgent — 仓位控制与风险预警。
根据市场波动率、回撤、相关性输出仓位建议和止损线。
"""
import numpy as np
import pandas as pd
from agents.base import BaseAgent
class RiskAgent(BaseAgent):
"""仓位控制 Agent。"""
name = "Risk"
description = "仓位控制与风险预警"
# 风险等级阈值
THRESHOLDS = {
"high": {"vol": 35, "dd": -15},
"medium": {"vol": 25, "dd": -8},
}
def execute(
self,
holdings: dict[str, float] | None = None,
market_index: str = "000001.SH", # 上证指数,非个股
ts_codes: list[str] | None = None,
) -> dict:
"""
评估市场风险并输出仓位建议。
参数:
holdings: {ts_code: 持仓比例}
market_index: 市场参考标的
ts_codes: 持仓股票列表
返回:
{"risk_level": str, "target_exposure": float, "indicators": dict, "alerts": list}
"""
holdings = holdings or {}
price = self.dm.get_daily(market_index)
if price is None or price.empty:
return self._default_result()
price = price.set_index("trade_date").sort_index()
close = price["close"]
daily_ret = close.pct_change().dropna()
# 市场波动率(年化)
market_vol = float(daily_ret.tail(252).std() * np.sqrt(252) * 100) if len(daily_ret) >= 20 else 30
# 当前回撤
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
# 最近 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)
# 持仓预警
alerts = []
if current_dd < -10:
alerts.append(f"市场回撤 {current_dd:.1f}%,考虑减仓")
if market_vol > 30:
alerts.append(f"市场波动率 {market_vol:.1f}%,处于高位")
for code, pct in holdings.items():
if pct > max_single:
alerts.append(f"{code} 仓位 {pct:.0%} 超过上限 {max_single:.0%}")
self.log(f"风险={risk_level} 波动={market_vol:.1f}% 回撤={current_dd:.1f}% 仓位→{target_exposure:.0%}")
return {
"risk_level": risk_level,
"target_exposure": round(target_exposure, 2),
"max_single_position": round(max_single, 2),
"stop_loss": round(stop_loss, 2),
"indicators": {
"market_volatility": round(market_vol, 1),
"current_drawdown": round(current_dd, 1),
"return_5d": round(ret_5d, 1),
"return_20d": round(ret_20d, 1),
"close": round(float(close.iloc[-1]), 2),
},
"alerts": alerts,
}
@staticmethod
def _default_result() -> dict:
return {
"risk_level": "medium",
"target_exposure": 0.60,
"max_single_position": 0.10,
"stop_loss": -0.08,
"indicators": {},
"alerts": ["数据不足,使用默认风险参数"],
}
+195
View File
@@ -0,0 +1,195 @@
"""
SelectionAgent — 多因子股票打分。
综合技术因子 + 基本面因子 + ML 预测值 → 股票综合得分排序。
"""
import numpy as np
import pandas as pd
from agents.base import BaseAgent
class SelectionAgent(BaseAgent):
"""股票打分 Agent。"""
name = "Selection"
description = "多因子股票打分与推荐"
def execute(
self,
date: str | None = None,
ts_codes: list[str] | None = None,
top_n: int = 15,
weighting: str = "equal",
) -> dict:
"""
对股票池打分排序。
参数:
date: 目标日期,None=最新
ts_codes: 股票池,None=指数成分股
top_n: 返回 Top N
weighting: 'equal' | 'ic_weighted' | 'ml'
返回:
{"date": ..., "top_picks": [...], "score_df": DataFrame}
"""
from factors.registry import get_factor, list_factors
# 确定股票池
if ts_codes is None:
if self.sent:
ts_codes = self.sent.get_scope_stocks()
else:
stocks = self.dm.get_stock_list()
ts_codes = list(stocks.index[:100])
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",
]
factor_objects = [get_factor(n) for n in core_factors]
self.log("打分 {} 只股票 (权重={})".format(len(ts_codes), weighting))
# 1. 筛选有 DB 缓存的股票(Orchestrator 已在 Step1 同步了全量范围)
from database.dao import get_latest_trade_date
available = []
for ts_code in sorted(ts_codes):
if get_latest_trade_date(ts_code):
available.append(ts_code)
self.log("缓存命中: {}/{} ({:.1f}%)".format(
len(available), len(ts_codes),
len(available) / len(ts_codes) * 100 if ts_codes else 0))
# 2. 打分(限制上限防止单次太慢)
score_limit = min(len(available), 300)
scores = {}
valid_count = 0
for i, ts_code in enumerate(available[:score_limit]):
try:
score = self._score_stock(ts_code, factor_objects, date, weighting)
if score is not None:
scores[ts_code] = score
valid_count += 1
except Exception:
continue
if (i + 1) % 50 == 0:
self.log(" 进度: {}/{}".format(i + 1, score_limit))
self.log("有效评分: {}/{}".format(len(scores), score_limit))
if not scores:
return {"date": date or self._today(), "top_picks": [], "score_df": pd.DataFrame()}
# 排序
sorted_scores = sorted(scores.items(), key=lambda x: x[1], reverse=True)
top = sorted_scores[:top_n]
# 取股票名称(ts_code 是 index
stock_list = self.dm.get_stock_list() if self.dm else pd.DataFrame()
name_map = {}
if not stock_list.empty and "name" in stock_list.columns:
name_map = dict(zip(stock_list.index, stock_list["name"]))
top_picks = []
for ts_code, score in top:
# 尝试多种格式匹配名称
code_no_suffix = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
name = name_map.get(ts_code, name_map.get(code_no_suffix, ts_code))
top_picks.append({
"ts_code": ts_code,
"name": name,
"score": round(score, 4),
})
score_df = pd.DataFrame(
{"ts_code": list(scores.keys()), "score": list(scores.values())}
).sort_values("score", ascending=False).reset_index(drop=True)
return {
"date": date or self._today(),
"top_picks": top_picks,
"score_df": score_df,
"universe_size": len(ts_codes),
"valid_scores": valid_count,
}
def _filter_cached_stocks(self, ts_codes: list[str], limit: int = 100) -> list[str]:
"""筛选有 DB 日线缓存的股票,避免逐个调用 AkShare。"""
from database.dao import get_latest_trade_date
cached = []
for ts_code in ts_codes[:limit]:
latest = get_latest_trade_date(ts_code)
if latest:
cached.append(ts_code)
return cached
def _score_stock(
self,
ts_code: str,
factors: list,
date: str | None,
weighting: str,
) -> float | None:
"""对单只股票打分。"""
daily = self.dm.get_daily(ts_code)
if daily is None or daily.empty:
return None
daily = daily.set_index("trade_date").sort_index()
factor_df = self.fe.compute(ts_code, factors)
if factor_df is None or factor_df.empty:
return None
if date and date in factor_df.index:
row = factor_df.loc[date]
else:
row = factor_df.iloc[-1] # 最新一天
if row.isna().all():
return None
if weighting == "ml" and self.ml_models:
return self._score_ml(ts_code, factor_df, date)
# 等权打分:标准化因子值后求和
row_clean = row.dropna()
if len(row_clean) < 3:
return None
# z-score 标准化
z = (row_clean - factor_df[row_clean.index].mean()) / factor_df[row_clean.index].std().replace(0, 1)
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)
daily = self.dm.get_daily(ts_code)
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:
return None
if date and date in X.index:
X = X.loc[[date]]
else:
X = X.iloc[[-1]]
model = self.ml_models.get("lightgbm") or list(self.ml_models.values())[0]
pred = model.predict(X)
return float(pred.iloc[0]) if len(pred) > 0 else None
except Exception:
return None
View File
+53
View File
@@ -0,0 +1,53 @@
"""
策略抽象基类。
所有策略必须继承 BaseStrategy,实现 generate_signals(factor_df) → pd.Series。
"""
from abc import ABC, abstractmethod
import pandas as pd
class BaseStrategy(ABC):
"""
回测策略基类。
属性:
name: 策略名称
category: 'trend' | 'mean_revert' | 'rotation'
"""
name: str = ""
category: str = ""
@abstractmethod
def generate_signals(self, factor_df: pd.DataFrame) -> pd.Series:
"""
因子 → 交易信号。
参数:
factor_df: 因子 DataFrameindex=trade_datecolumns=因子名
返回:
pd.Seriesindex 与 factor_df 对齐:
1=买入, 0=平仓/无操作
(只做多,不做空)
"""
...
def get_params(self) -> dict:
"""返回策略当前参数(供 Optuna 优化用)。"""
return {
k: v for k, v in self.__dict__.items()
if not k.startswith("_") and k not in ("name", "category")
}
def set_params(self, **kwargs) -> None:
"""设置策略参数。"""
for k, v in kwargs.items():
if hasattr(self, k):
setattr(self, k, v)
def __repr__(self) -> str:
return f"{self.__class__.__name__}(name='{self.name}')"
+144
View File
@@ -0,0 +1,144 @@
"""
标准化回测报告。
与回测引擎解耦,后续换引擎只需改构造函数。
"""
from dataclasses import dataclass, field
import numpy as np
import pandas as pd
@dataclass
class BacktestReport:
"""回测报告"""
# 核心收益指标
total_return: float = 0.0
cagr: float = 0.0
max_drawdown: float = 0.0
sharpe_ratio: float = 0.0
calmar_ratio: float = 0.0
annual_volatility: float = 0.0
# 交易统计
win_rate: float = 0.0
profit_factor: float = 0.0
total_trades: int = 0
avg_hold_days: float = 0.0
best_trade_pct: float = 0.0
worst_trade_pct: float = 0.0
# 序列数据
equity_curve: pd.Series = field(default_factory=pd.Series)
drawdown_curve: pd.Series = field(default_factory=pd.Series)
monthly_returns: pd.Series = field(default_factory=pd.Series)
trades_df: pd.DataFrame = field(default_factory=pd.DataFrame)
# 原始 stats
stats_dict: dict = field(default_factory=dict)
@classmethod
def from_vbt_result(cls, pf, close: pd.Series) -> "BacktestReport":
"""从 VectorBT Portfolio 结果构建报告。"""
stats = pf.stats()
def _pct(v):
"""VectorBT stats 值已为百分比 float,直接返回。"""
if v is None:
return 0.0
try:
return float(v)
except (ValueError, TypeError):
return 0.0
def _duration_days(v):
"""Timedelta → 天数。"""
if v is None:
return 0.0
try:
return v.total_seconds() / 86400
except AttributeError:
return float(v) if v is not None else 0.0
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")
dd = equity / equity.cummax() - 1
daily_ret = equity.pct_change().dropna()
years = len(daily_ret) / 252 if len(daily_ret) > 0 else 1
total_ret = (equity.iloc[-1] / equity.iloc[0] - 1) * 100 if len(equity) > 1 else 0
cagr = ((total_ret / 100 + 1) ** (1 / years) - 1) * 100 if years > 0 else 0
mdd = dd.min() * 100
ann_vol = daily_ret.std() * np.sqrt(252) * 100 if len(daily_ret) > 0 else 0
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 mdd != 0 else 0
trades = pf.trades.records_readable if hasattr(pf, "trades") else pd.DataFrame()
try:
monthly = equity.resample("ME").last().pct_change()
except Exception:
monthly = pd.Series(dtype=float)
pf_factor = stats.get("Profit Factor", 0)
if pf_factor is None or np.isinf(float(pf_factor)):
pf_factor = 0.0
else:
pf_factor = float(pf_factor)
return cls(
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(ann_vol, 2),
win_rate=_pct(stats.get("Win Rate [%]", 0)),
profit_factor=pf_factor,
total_trades=int(stats.get("Total Trades", 0)),
avg_hold_days=_duration_days(stats.get("Avg Winning Trade Duration", None)),
best_trade_pct=_pct(stats.get("Best Trade [%]", 0)),
worst_trade_pct=_pct(stats.get("Worst Trade [%]", 0)),
equity_curve=equity,
drawdown_curve=dd,
monthly_returns=monthly,
trades_df=trades,
stats_dict={k: str(v) for k, v in stats.items()},
)
def to_dict(self) -> dict:
"""核心指标转字典。"""
return {
"total_return": self.total_return,
"cagr": self.cagr,
"max_drawdown": self.max_drawdown,
"sharpe_ratio": self.sharpe_ratio,
"calmar_ratio": self.calmar_ratio,
"annual_volatility": self.annual_volatility,
"win_rate": self.win_rate,
"profit_factor": self.profit_factor,
"total_trades": self.total_trades,
}
def summary(self) -> str:
"""一行摘要。"""
return (
f"收益={self.total_return:.1f}% "
f"年化={self.cagr:.1f}% "
f"回撤={self.max_drawdown:.1f}% "
f"夏普={self.sharpe_ratio:.2f} "
f"胜率={self.win_rate:.1f}% "
f"交易={self.total_trades}"
)
def __repr__(self) -> str:
return self.summary()
+115
View File
@@ -0,0 +1,115 @@
"""
信号生成工具函数。
因子值 → 交易信号的桥梁,纯函数无副作用。
"""
import pandas as pd
def factor_to_threshold_signal(
factor_series: pd.Series,
buy_threshold: float,
sell_threshold: float | None = None,
cross_direction: str = "up",
) -> pd.Series:
"""
因子阈值交叉信号。
参数:
factor_series: 因子值 Series
buy_threshold: 买入阈值(如 RSI < 30 则买)
sell_threshold: 卖出阈值(如 RSI > 70 则卖),None 表示平所有仓
cross_direction: 'up'=因子向上穿越阈值时触发, 'down'=向下穿越
返回:
信号 Series1=买入, 0=平仓
"""
signals = pd.Series(0, index=factor_series.index)
if cross_direction == "down":
buys = factor_series < buy_threshold
else:
buys = factor_series > buy_threshold
signals[buys] = 1
if sell_threshold is not None:
if cross_direction == "down":
sells = factor_series > sell_threshold
else:
sells = factor_series < sell_threshold
signals[sells] = 0
# 过滤连续信号
signals = _filter_consecutive(signals)
return signals
def factor_to_quantile_signal(
factor_series: pd.Series,
top_quantile: float = 0.8,
bottom_quantile: float = 0.2,
) -> pd.Series:
"""
因子分位数信号 — 按滚动分位数判断。
参数:
factor_series: 因子值
top_quantile: 高于此分位买入
bottom_quantile: 低于此分位平仓
返回:
信号 Series
"""
top = factor_series.quantile(top_quantile)
bottom = factor_series.quantile(bottom_quantile)
signals = pd.Series(0, index=factor_series.index)
signals[factor_series > top] = 1
signals[factor_series < bottom] = 0
return _filter_consecutive(signals)
def cross_signal(
fast: pd.Series,
slow: pd.Series,
) -> pd.Series:
"""
金叉/死叉信号。
fast 上穿 slow → 买入(1)
fast 下穿 slow → 平仓(0)
"""
fast = fast.dropna()
slow = slow.dropna()
common_idx = fast.index.intersection(slow.index)
fast, slow = fast[common_idx], slow[common_idx]
signals = pd.Series(-1, index=common_idx)
above = (fast > slow).fillna(False)
above = above.infer_objects(copy=False)
# 交叉点:今天 above=True 且昨天 above=False → 金叉
prev = above.shift(1).fillna(False)
prev = prev.infer_objects(copy=False)
cross_up = above & ~prev
cross_down = ~above & prev
signals[cross_up] = 1
signals[cross_down] = 0
return _filter_consecutive(signals)
def _filter_consecutive(signals: pd.Series) -> pd.Series:
"""过滤连续相同信号,只保留首次出现的信号。"""
result = signals.copy()
prev = None
for i in range(len(result)):
if result.iloc[i] == prev:
result.iloc[i] = -1 # 标记为不操作
else:
prev = result.iloc[i]
return result[result != -1].reindex(signals.index).fillna(-1)
@@ -0,0 +1,51 @@
"""
因子阈值交叉策略。
通用策略:任意因子上穿/下穿阈值 → 交易信号。
支持:
- 上穿买入 (cross_up: close < MA → cross above MA → buy)
- 下穿买入 (cross_down: RSI > 70 → cross below 30 → buy)
"""
import pandas as pd
from backtest.base import BaseStrategy
from backtest.signal import factor_to_threshold_signal
class FactorCrossStrategy(BaseStrategy):
"""
因子阈值交叉策略。
适用场景:
- 均线偏离度上穿 0 → 买入(趋势转多)
- 波动率下穿阈值 → 买入(波动收敛后突破)
"""
category = "trend"
def __init__(
self,
factor_column: str,
buy_threshold: float, # 因子大于此值买
sell_threshold: float | None = None,
cross_direction: str = "up",
):
self.factor_column = factor_column
self.buy_threshold = buy_threshold
self.sell_threshold = sell_threshold
self.cross_direction = cross_direction
self.name = f"factor_cross_{factor_column}"
def generate_signals(self, factor_df: pd.DataFrame) -> pd.Series:
if self.factor_column not in factor_df.columns:
raise ValueError(f"factor_df 缺少 '{self.factor_column}'")
factor = factor_df[self.factor_column]
return factor_to_threshold_signal(
factor,
buy_threshold=self.buy_threshold,
sell_threshold=self.sell_threshold,
cross_direction=self.cross_direction,
)
@@ -0,0 +1,58 @@
"""
因子排序轮动策略。
定期按因子值排序,买入排名最高的股票(截面策略)。
"""
import pandas as pd
from backtest.base import BaseStrategy
class FactorRotationStrategy(BaseStrategy):
"""
因子排序选股策略。
适用于多股票截面场景:对每只股票计算因子值,
选排名最高的 top_n 只做多。
"""
category = "rotation"
def __init__(self, factor_name: str, top_n: int = 5, bottom_n: int = 0):
self.factor_name = factor_name
self.top_n = top_n
self.bottom_n = bottom_n
self.name = f"rotation_{factor_name}_top{top_n}"
def generate_signals(self, factor_df: pd.DataFrame) -> pd.Series:
"""
单股票/截面模式:factor_df 支持两种输入方式。
- 单股票: 对每只股票单次调用
- 截面: 通过 run_cross_section 逐股票调用
"""
if self.factor_name not in factor_df.columns:
raise ValueError(f"factor_df 缺少 '{self.factor_name}'")
factor = factor_df[self.factor_name]
valid = factor.dropna()
if len(valid) < self.top_n * 2:
return pd.Series(-1, index=factor_df.index)
threshold = valid.quantile(1 - self.top_n / max(len(valid), self.top_n))
signals = pd.Series(-1, index=factor_df.index)
signals[factor > threshold] = 1
return signals
def rank_stocks(
self, factor_values: dict[str, float]
) -> list[str]:
"""
对股票按因子值排序,返回 top N 的 ts_code 列表。
参数:
factor_values: {ts_code: factor_value}
"""
sorted_stocks = sorted(factor_values, key=factor_values.get, reverse=True)
return sorted_stocks[: self.top_n]
@@ -0,0 +1,52 @@
"""
动量突破策略。
价格突破 N 日新高 → 买入
价格跌破 N 日均线 → 平仓
"""
import pandas as pd
from backtest.base import BaseStrategy
from backtest.signal import factor_to_threshold_signal
class MomentumBreakoutStrategy(BaseStrategy):
"""动量突破策略。"""
category = "trend"
def __init__(self, lookback: int = 20, exit_period: int = 10):
self.lookback = lookback
self.exit_period = exit_period
self.name = f"mom_breakout_{lookback}"
def generate_signals(self, factor_df: pd.DataFrame) -> pd.Series:
if "close" not in factor_df.columns:
raise ValueError("factor_df 缺少 'close'")
close = factor_df["close"]
# 买入信号:突破 N 日新高
rolling_high = close.rolling(window=self.lookback, min_periods=self.lookback).max()
breakout = close >= rolling_high.shift(1)
# 平仓信号:跌破 exit 日均线
exit_ma = close.rolling(window=self.exit_period, min_periods=self.exit_period).mean()
signals = pd.Series(0, index=close.index)
signals[breakout] = 1
signals[close < exit_ma] = 0
return self._dedup(signals)
@staticmethod
def _dedup(signals: pd.Series) -> pd.Series:
"""只保留第一个买入和第一个卖出信号。"""
result = signals.copy()
prev = -1
for i in range(len(result)):
if result.iloc[i] == prev:
result.iloc[i] = -1
else:
prev = result.iloc[i]
return result[result != -1].reindex(signals.index).fillna(-1)
@@ -0,0 +1,35 @@
"""
RSI 均值回归策略。
RSI 低于超卖线 → 买入
RSI 高于超买线 → 平仓
"""
import pandas as pd
from backtest.base import BaseStrategy
from backtest.signal import factor_to_threshold_signal
class RSIMeanRevertStrategy(BaseStrategy):
"""RSI 超买超卖反转策略。"""
category = "mean_revert"
def __init__(self, oversold: float = 30, overbought: float = 70, rsi_column: str = "rsi_14"):
self.oversold = oversold
self.overbought = overbought
self.rsi_column = rsi_column
self.name = f"rsi_revert_{int(oversold)}_{int(overbought)}"
def generate_signals(self, factor_df: pd.DataFrame) -> pd.Series:
if self.rsi_column not in factor_df.columns:
raise ValueError(f"factor_df 缺少 '{self.rsi_column}'")
rsi = factor_df[self.rsi_column]
return factor_to_threshold_signal(
rsi,
buy_threshold=self.oversold,
sell_threshold=self.overbought,
cross_direction="down", # RSI 向下跌破 oversold → 买入
)
+33
View File
@@ -0,0 +1,33 @@
"""
均线交叉策略。
短期均线上穿长期均线 → 买入
短期均线下穿长期均线 → 平仓
"""
import pandas as pd
from backtest.base import BaseStrategy
from backtest.signal import cross_signal
class SMACrossStrategy(BaseStrategy):
"""快慢均线交叉策略。"""
category = "trend"
def __init__(self, fast: int = 5, slow: int = 20):
self.fast = fast
self.slow = slow
self.name = f"sma_cross_{fast}_{slow}"
def generate_signals(self, factor_df: pd.DataFrame) -> pd.Series:
if "close" not in factor_df.columns:
raise ValueError("factor_df 缺少 'close'")
close = factor_df["close"]
min_p = min(self.fast, self.slow)
ma_fast = close.rolling(window=self.fast, min_periods=self.fast).mean()
ma_slow = close.rolling(window=self.slow, min_periods=self.slow).mean()
return cross_signal(ma_fast, ma_slow)
+223
View File
@@ -0,0 +1,223 @@
"""
VectorBT 回测引擎封装。
统一接口:engine.run(strategy, price_df, factor_df) → BacktestReport
"""
import numpy as np
import pandas as pd
import vectorbt as vbt
from backtest.base import BaseStrategy
from backtest.report import BacktestReport
class VectorBTEngine:
"""
VectorBT 回测引擎。
只做多,不做空。
"""
def __init__(
self,
initial_capital: float = 100_000,
commission: float = 0.0003, # 万三
freq: str = "D",
):
self.initial_capital = initial_capital
self.commission = commission
self.freq = freq
# ── 单股票回测 ────────────────────────────────────────
def run(
self,
strategy: BaseStrategy,
price_df: pd.DataFrame,
factor_df: pd.DataFrame | None = None,
) -> BacktestReport:
"""
单股票回测。
参数:
strategy: 策略实例
price_df: 价格数据,index=trade_date,必须有 'close'
factor_df: 因子数据,index=trade_date。
None 时使用 price_df 作为因子数据源。
返回:
BacktestReport
"""
if factor_df is None:
factor_df = price_df
# 1. 对齐日期
common_idx = price_df.index.intersection(factor_df.index)
if len(common_idx) < 2:
return BacktestReport()
price_df = price_df.loc[common_idx].sort_index()
factor_df = factor_df.loc[common_idx].sort_index()
# 2. 合并 close 到 factor_df(策略可能需要)
if "close" not in factor_df.columns:
factor_df = factor_df.copy()
factor_df["close"] = price_df["close"]
# 3. 生成信号
raw_signals = strategy.generate_signals(factor_df)
# 4. 信号 → VectorBT entries/exits
entries, exits = self._signals_to_entries(raw_signals, price_df.index)
# 5. 运行回测
close = price_df["close"]
pf = vbt.Portfolio.from_signals(
close,
entries=entries,
exits=exits,
init_cash=self.initial_capital,
fees=self.commission,
freq=self.freq,
direction="longonly",
)
return BacktestReport.from_vbt_result(pf, close)
# ── 截面回测(多股票) ──────────────────────────────────
def run_cross_section(
self,
strategy: BaseStrategy,
price_universe: dict[str, pd.DataFrame],
factor_universe: dict[str, pd.DataFrame] | None = None,
rebalance_freq: str = "M",
) -> BacktestReport:
"""
截面策略回测(多股票 + 定期调仓)。
对每只股票独立回测,合并权益曲线。
参数:
strategy: 策略实例
price_universe: {ts_code: price_df}
factor_universe: {ts_code: factor_df}
rebalance_freq: 调仓频率 'D'/'W'/'M',用于合并时对齐
返回:
BacktestReport
"""
if factor_universe is None:
factor_universe = price_universe
stock_equities = {}
stock_reports = {}
# 逐股票回测
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)
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
if not stock_equities:
return BacktestReport()
# 合并:等权分配资金到各股票
return self._merge_equities(stock_equities)
# ── 信号转换 ──────────────────────────────────────────
@staticmethod
def _signals_to_entries(
raw_signals: pd.Series,
target_index: pd.Index,
) -> tuple[pd.Series, pd.Series]:
"""
将策略信号转为 VectorBT entries/exits。
信号格式:
1 → 买入
0 → 平仓
-1 → 继续持有/不操作
entries: True 时开仓
exits: True 时平仓
"""
# 对齐到目标 index
aligned = pd.Series(-1, index=target_index)
common = target_index.intersection(raw_signals.index)
aligned.loc[common] = raw_signals.loc[common].values
entries = pd.Series(False, index=target_index)
exits = pd.Series(False, index=target_index)
in_position = False
for i in range(len(aligned)):
sig = aligned.iloc[i]
if not in_position and sig == 1:
entries.iloc[i] = True
in_position = True
elif in_position and sig == 0:
exits.iloc[i] = True
in_position = False
return entries, exits
# ── 合并多股票权益 ─────────────────────────────────────
def _merge_equities(
self, stock_equities: dict[str, pd.Series]
) -> BacktestReport:
"""等权合并多股票权益曲线,构建组合级报告。"""
equity_df = pd.DataFrame(stock_equities)
equity_df = equity_df.ffill().fillna(0)
# 转为 DatetimeIndex
if not isinstance(equity_df.index, pd.DatetimeIndex):
equity_df.index = pd.to_datetime(equity_df.index, format="%Y%m%d")
n_stocks = len(stock_equities)
weight = 1.0 / n_stocks if n_stocks > 0 else 1.0
# 加权组合收益
returns_df = equity_df.pct_change().fillna(0)
portfolio_ret = returns_df.mean(axis=1) # 等权 = 逐行平均
# 组合净值
portfolio_equity = self.initial_capital * (1 + portfolio_ret).cumprod()
dd = portfolio_equity / portfolio_equity.cummax() - 1
years = max(len(portfolio_ret) / 252, 0.02)
total_return = (portfolio_equity.iloc[-1] / portfolio_equity.iloc[0] - 1) * 100
cagr = ((total_return / 100 + 1) ** (1 / years) - 1) * 100
mdd = dd.min() * 100
mean_ret = portfolio_ret.mean() * 252
std_ret = portfolio_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_return, 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,
)
View File
+152
View File
@@ -0,0 +1,152 @@
#!/usr/bin/env python
"""
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] # 生成日报
"""
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from datetime import datetime
def init_engines():
"""初始化所有引擎。"""
from data.data_manager import DataManager
from factors.engine import FactorEngine
from backtest.vectorbt.engine import VectorBTEngine
from optimizer.engine import OptunaEngine
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()
fe = FactorEngine(dm)
bt = VectorBTEngine()
opt = OptunaEngine(bt)
sent = SentimentEngine(dm, qwen_client=QwenClient(), news_source=NewsSource())
fe._sentiment_engine = sent
return {
"dm": dm,
"fe": fe,
"bt": bt,
"opt": opt,
"sent": sent,
}
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
cmd = sys.argv[1]
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)
# 打印日报内容
report_md = results.get("report", {}).get("report_markdown", "")
if report_md:
print(report_md)
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)
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}")
elif cmd == "risk":
result = orch.risk_check()
print(f"\n风险评估:")
print(f" 等级: {result['risk_level']}")
print(f" 建议仓位: {result['target_exposure']:.0%}")
print(f" 止损线: {result['stop_loss']:.0%}")
print(f" 单票上限: {result['max_single_position']:.0%}")
indicators = result.get("indicators", {})
if indicators:
print(f" 波动率: {indicators.get('market_volatility', 0):.1f}%")
print(f" 回撤: {indicators.get('current_drawdown', 0):.1f}%")
for a in result.get("alerts", []):
print(f" ⚠️ {a}")
elif cmd == "research":
result = orch.run_research_cycle()
top = result.get("research", {}).get("top_factors", [])
print(f"\n因子评估结果:")
if not top:
print(" (无结果)")
return
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}")
elif cmd == "warmup":
batch_n = int(sys.argv[2]) if len(sys.argv) > 2 else 50
print("首次批量预热: 每次 {} 只股票,分批执行...".format(batch_n))
sent = engines.get("sent")
dm = engines.get("dm")
scope = sent.get_scope_stocks() if sent else list(dm.get_stock_list().index[:100])
from database.dao import get_latest_trade_date
uncached = [c for c in scope if not get_latest_trade_date(c)]
print("范围: {} 只, 未缓存: {}".format(len(scope), len(uncached)))
total_synced = 0
for i in range(0, len(uncached), batch_n):
batch = uncached[i:i + batch_n]
print("[warmup] 批次 {}/{} ({}~{})".format(i // batch_n + 1, (len(uncached) - 1) // batch_n + 1, i, i + len(batch)))
for ts_code in batch:
try:
n = dm.sync_daily(ts_code)
total_synced += n
except Exception as e:
print(" {} 失败: {}".format(ts_code, e))
print(" 累计同步: {}".format(total_synced))
print("预热完成: {} 条数据, {} 只新股票已缓存".format(total_synced, len(uncached)))
elif cmd == "report":
date = sys.argv[2] if len(sys.argv) > 2 else None
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")
if md:
print(md)
else:
print(f"未知命令: {cmd}")
if __name__ == "__main__":
main()
+89
View File
@@ -0,0 +1,89 @@
"""
Sprint 2 验证脚本 — 回测引擎。
用法:
python cli/demo_backtest.py
python cli/demo_backtest.py --ts_code 600519.SH
"""
import sys, os, argparse
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from data.data_manager import DataManager
from factors.engine import FactorEngine
from factors.registry import get_factor
from backtest.vectorbt.engine import VectorBTEngine
from backtest.strategies.sma_cross import SMACrossStrategy
from backtest.strategies.rsi_mean_revert import RSIMeanRevertStrategy
from backtest.strategies.momentum_breakout import MomentumBreakoutStrategy
from backtest.strategies.factor_cross import FactorCrossStrategy
from backtest.strategies.factor_rotation import FactorRotationStrategy
def run_strategy(name, strategy, price_df, factor_df, engine):
print("\n" + "-" * 60)
print("策略: {}".format(name))
try:
report = engine.run(strategy, price_df, factor_df)
print(report.summary())
return report
except Exception as e:
print(" [FAIL] {}".format(e))
return None
def main():
p = argparse.ArgumentParser(description="Sprint 2 — 回测引擎验证")
p.add_argument("--ts_code", default="000001.SZ", help="测试股票代码(默认: 000001.SZ)")
args = p.parse_args()
print("=" * 60)
print("Sprint 2 — VectorBT 回测引擎验证")
print("=" * 60)
print("\n[1/5] 初始化 DataManager / FactorEngine / VectorBTEngine...")
dm = DataManager(); dm.init_db()
engine_fe = FactorEngine(dm)
engine_bt = VectorBTEngine()
price_df = dm.get_daily(args.ts_code).set_index("trade_date").sort_index()
print(" 日线: {} 条 ({} ~ {})".format(len(price_df), price_df.index[0], price_df.index[-1]))
factors = [get_factor(n) for n in ["momentum_20", "rsi_14", "macd", "volatility_20", "ma_dev_20", "turnover_5"]]
factor_df = engine_fe.compute(args.ts_code, factors)
print(" 因子: {} 个, {} 个交易日".format(factor_df.shape[1], factor_df.shape[0]))
print("[OK] 就绪")
results = []
print("\n[2/5] 测试均线交叉策略...")
results.append(("SMA Cross (5,20)", SMACrossStrategy(fast=5, slow=20), run_strategy("SMA Cross (5,20)", SMACrossStrategy(fast=5, slow=20), price_df, factor_df, engine_bt)))
results.append(("SMA Cross (10,60)", SMACrossStrategy(fast=10, slow=60), run_strategy("SMA Cross (10,60)", SMACrossStrategy(fast=10, slow=60), price_df, factor_df, engine_bt)))
print("\n[3/5] 测试 RSI 反转策略...")
results.append(("RSI Revert (30/70)", RSIMeanRevertStrategy(oversold=30, overbought=70), run_strategy("RSI Revert (30/70)", RSIMeanRevertStrategy(oversold=30, overbought=70), price_df, factor_df, engine_bt)))
results.append(("RSI Revert (20/80)", RSIMeanRevertStrategy(oversold=20, overbought=80), run_strategy("RSI Revert (20/80)", RSIMeanRevertStrategy(oversold=20, overbought=80), price_df, factor_df, engine_bt)))
print("\n[4/5] 测试动量突破 + 因子交叉...")
results.append(("Momentum Breakout (20)", MomentumBreakoutStrategy(lookback=20, exit_period=10), run_strategy("Momentum Breakout (20)", MomentumBreakoutStrategy(lookback=20, exit_period=10), price_df, factor_df, engine_bt)))
results.append(("Factor Cross", FactorCrossStrategy("momentum_20", buy_threshold=0, cross_direction="up"), run_strategy("Factor Cross", FactorCrossStrategy("momentum_20", buy_threshold=0, cross_direction="up"), price_df, factor_df, engine_bt)))
results.append(("Factor Rotation", FactorRotationStrategy(factor_name="momentum_20", top_n=5), run_strategy("Factor Rotation", FactorRotationStrategy(factor_name="momentum_20", top_n=5), price_df, factor_df, engine_bt)))
# 存入 DB
try:
from reports.storage import save_report
summary = "## 回测验证 — {}\n\n".format(args.ts_code)
for name, s, report in results:
if report:
summary += "### {}\n{}\n\n".format(name, report.summary())
save_report(summary, "回测验证", subject_type="stock", subject_code=args.ts_code)
print("\n 报告已存入 DB")
except Exception as e:
print("\n [WARN] 报告入库失败: {}".format(e))
print("\n" + "=" * 60)
print("Sprint 2 验证完成")
print("=" * 60)
if __name__ == "__main__":
main()
+76
View File
@@ -0,0 +1,76 @@
"""
Sprint 0 验证脚本 — DataManager 全链路。
用法:
python cli/demo_data_manager.py
python cli/demo_data_manager.py --ts_code 600519.SH
python cli/demo_data_manager.py --ts_code 300316.SZ --start 20250101
"""
import sys, os, argparse
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from database.connection import test_connection
from database.models import create_all_tables
from data.data_manager import DataManager
def main():
p = argparse.ArgumentParser(description="Sprint 0 — DataManager 验证")
p.add_argument("--ts_code", default="000001.SZ", help="测试股票代码(默认: 000001.SZ)")
p.add_argument("--start", default="20250101", help="起始日期 YYYYMMDD(默认: 20250101")
args = p.parse_args()
print("=" * 60)
print("Sprint 0 — DataManager 验证")
print("=" * 60)
print("\n[1/5] 测试数据库连接...")
if not test_connection():
print("请先执行 shared/script/autossh.sh 建立 SSH 隧道")
return
print("[OK] 数据库连接成功")
print("\n[2/5] 创建数据表...")
dm = DataManager(); dm.init_db()
print("[OK] 表结构已就绪")
print("\n[3/5] 获取股票列表...")
stocks = dm.get_stock_list()
print("[OK] 共 {} 只股票".format(len(stocks)))
print(stocks.head(10))
print("\n[4/5] 获取日线数据 ({} start={})...".format(args.ts_code, args.start))
try:
daily = dm.get_daily(args.ts_code, start=args.start)
if not daily.empty:
print("[OK] 获取到 {} 条日线".format(len(daily)))
print(" 日期范围: {} ~ {}".format(daily['trade_date'].min(), daily['trade_date'].max()))
print(daily.tail(5))
else:
print("[WARN] 日线数据为空(AkShare+Tushare 均不可用)")
except Exception as e:
print("[WARN] 日线获取异常: {}".format(e))
print("\n[5/5] 增量同步测试...")
try:
count = dm.sync_daily(args.ts_code)
print("[OK] 增量同步结果: {}".format(count))
except Exception as e:
print("[WARN] 增量同步异常: {}".format(e))
try:
from reports.storage import save_report
save_report("## 数据层验证 — {}\n\n- 股票列表: OK\n- 日线数据: OK\n- 增量同步: OK".format(args.ts_code),
"数据层验证", subject_type="stock", subject_code=args.ts_code)
print("\n 报告已存入 DB")
except Exception:
pass
print("\n" + "=" * 60)
print("验证完成")
print("=" * 60)
if __name__ == "__main__":
main()
+85
View File
@@ -0,0 +1,85 @@
"""
Sprint 1 验证脚本 — 因子引擎。
用法:
python cli/demo_factor_engine.py
python cli/demo_factor_engine.py --ts_code 600519.SH
"""
import sys, os, argparse
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import pandas as pd
from data.data_manager import DataManager
from factors.registry import get_factor, list_factors, list_categories
from factors.engine import FactorEngine
def main():
p = argparse.ArgumentParser(description="Sprint 1 — 因子引擎验证")
p.add_argument("--ts_code", default="000001.SZ", help="测试股票代码(默认: 000001.SZ)")
p.add_argument("--ts_code2", default="600519.SH", help="截面测试第二只股票(默认: 600519.SH)")
args = p.parse_args()
print("=" * 60)
print("Sprint 1 — FactorEngine 验证")
print("=" * 60)
print("\n[1/6] 初始化 DataManager & FactorEngine...")
dm = DataManager(); dm.init_db()
engine = FactorEngine(dm)
dm.get_stock_list()
_ = dm.get_daily(args.ts_code)
_ = dm.get_daily(args.ts_code2)
print("[OK] 就绪")
print("\n[2/6] 因子注册表...")
cats = list_categories()
print(" {} 个分类, {} 个因子".format(len(cats), len(list_factors())))
for cat in cats:
print(" [{}]: {}".format(cat, ", ".join(list_factors(cat))))
print("\n[3/6] 计算技术因子 ({})...".format(args.ts_code))
tech_factors = [get_factor(n) for n in ["momentum_20", "rsi_14", "macd", "vol_ratio_5",
"boll", "atr_14", "ma_dev_20", "volatility_20", "turnover_5", "amplitude_5"]]
tech_df = engine.compute(args.ts_code, tech_factors)
print(" shape: {}".format(tech_df.shape))
print(tech_df.describe().round(2).to_string())
print("\n[4/6] 计算基本面因子 ({},需财务数据)...".format(args.ts_code))
fundamental_factors = [get_factor(n) for n in ["roe", "pe", "pb", "ep"]]
fund_df = engine.compute(args.ts_code, fundamental_factors)
if fund_df is not None and not fund_df.empty:
valid = fund_df.dropna(how="all")
print(" 有效行: {}/{}".format(len(valid), len(fund_df)))
if not valid.empty:
print(valid.tail(5).round(2).to_string())
print("\n[5/6] 因子 NaN 覆盖率检查...")
all_df = engine.compute(args.ts_code, tech_factors + fundamental_factors)
for col in all_df.columns:
nan_pct = all_df[col].isna().sum() / len(all_df) * 100
print(" {:20s}: NaN {:5.1f}%".format(col, nan_pct))
print("\n[6/6] 截面因子 ({} + {})...".format(args.ts_code, args.ts_code2))
cross = engine.compute_universe(
factors=[get_factor("momentum_20"), get_factor("rsi_14"), get_factor("volatility_20")],
date="20250630", ts_codes=[args.ts_code, args.ts_code2],
)
print(cross.round(4).to_string() if not cross.empty else " (空)")
try:
from reports.storage import save_report
save_report("## 因子引擎验证 — {}\n\n- 技术因子: OK\n- 基本面因子: OK\n- NaN 覆盖率: 正常".format(args.ts_code),
"因子引擎验证", subject_type="stock", subject_code=args.ts_code)
print("\n 报告已存入 DB")
except Exception:
pass
print("\n" + "=" * 60)
print("Sprint 1 验证完成")
print("=" * 60)
if __name__ == "__main__":
main()
+175
View File
@@ -0,0 +1,175 @@
"""
Sprint 4 验证脚本 — ML 模型训练与回测。
用法:
python cli/demo_ml.py
python cli/demo_ml.py --ts_code 600519.SH
python cli/demo_ml.py --ts_code 000001.SZ --lookahead 10
"""
import sys, os, argparse
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import numpy as np
from data.data_manager import DataManager
from factors.engine import FactorEngine
from factors.registry import get_factor
from models.features import FeatureEngine
from models.lightgbm.model import LightGBMModel
from models.catboost.model import CatBoostModel
from models.backtest_integration import MLStrategy, MLBenchmark
from backtest.vectorbt.engine import VectorBTEngine
def main():
p = argparse.ArgumentParser(description="Sprint 4 — ML 模型训练与回测验证")
p.add_argument("--ts_code", default="000001.SZ", help="测试股票代码")
p.add_argument("--lookahead", type=int, default=5, help="预测未来 N 日(默认: 5")
args = p.parse_args()
print("=" * 60)
print("Sprint 4 — ML 模型训练与回测验证")
print("=" * 60)
print("\n[1/6] 准备数据...")
dm = DataManager(); dm.init_db()
engine_fe = FactorEngine(dm)
price_df = dm.get_daily(args.ts_code).set_index("trade_date").sort_index()
tech_names = ["momentum_5", "momentum_10", "momentum_20", "momentum_60", "rsi_7", "rsi_14", "macd",
"vol_ratio_5", "vol_ratio_20", "vol_chg_5", "boll", "boll_width", "atr_14", "atr_ratio_14",
"ma_cross_5_20", "ma_cross_10_60", "ma_dev_20", "ma_dev_60", "volatility_20", "volatility_60",
"down_vol_20", "turnover_5", "turnover_chg_5", "amplitude_5", "amplitude_20"]
factor_df = engine_fe.compute(args.ts_code, [get_factor(n) for n in tech_names])
print(" 日线: {} 条, 因子: {}".format(len(price_df), factor_df.shape[1]))
print("\n[2/6] 特征工程 (lookahead={}, regression)...".format(args.lookahead))
fe = FeatureEngine(lookahead=args.lookahead, label_type="regression")
X, y = fe.build(factor_df, price_df, fit=True)
print(" 特征矩阵: {} x {}".format(X.shape[0], X.shape[1]))
print(" 标签: mean={:.2f}% std={:.2f}% min={:.1f}% max={:.1f}%".format(y.mean(), y.std(), y.min(), y.max()))
print("\n[3/6] 训练集/测试集划分 (前70%训练)...")
n = len(X); split = int(n * 0.7)
X_train, X_test = X.iloc[:split], X.iloc[split:]
y_train, y_test = y.iloc[:split], y.iloc[split:]
print(" 训练集: {} 行 ({} ~ {})".format(len(X_train), X_train.index[0], X_train.index[-1]))
print(" 测试集: {} 行 ({} ~ {})".format(len(X_test), X_test.index[0], X_test.index[-1]))
print("\n[4/6] LightGBM 训练...")
lgb_model = LightGBMModel(eval_ratio=0.0, early_stopping=100)
lgb_model.fit(X_train, y_train)
ic_lgb = lgb_model.predict(X_test).corr(y_test)
print(" 测试集 IC: {:.4f} trees: {}".format(ic_lgb, lgb_model.n_estimators_used))
imp = lgb_model.get_feature_importance()
print(" Top 5 特征:")
for _, row in imp.head(5).iterrows():
print(" {:25s} {:5.1f}%".format(row["feature"], row["importance_pct"]))
# 解读
abs_ic = abs(ic_lgb)
if abs_ic < 0.03:
print(" > 解读: IC 接近 0,单股票预测信号极弱(正常现象)。多股票截面预测效果更好。")
elif abs_ic < 0.08:
print(" > 解读: IC {:.3f} 有微弱预测能力,可用于因子组合。".format(ic_lgb))
else:
print(" > 解读: IC {:.3f} 有显著预测能力,特征工程有效。".format(ic_lgb))
top_feat = imp.iloc[0]
print(" > 最重要特征 '{}' 占比 {:.1f}%,说明该类因子对短期收益影响最大。".format(top_feat["feature"], top_feat["importance_pct"]))
print("\n[5/6] CatBoost 训练...")
cb_model = CatBoostModel(eval_ratio=0.0, early_stopping=100)
cb_model.fit(X_train, y_train)
ic_cb = cb_model.predict(X_test).corr(y_test)
print(" 测试集 IC: {:.4f} trees: {}".format(ic_cb, cb_model.n_estimators_used))
if abs(ic_cb) < 0.03:
print(" > 解读: CatBoost IC 同样接近 0,两者结论一致:单股票短期收益很难预测。")
print("\n[6/6] ML 策略回测对比...")
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)
result = benchmark.run()
print(result.round(2).to_string())
print(" > 解读: 回测结果反映 ML 策略在测试集上的实盘表现。")
print(" > 正收益+高夏普=模型有效;负收益=需更多特征或换截面预测。")
print(" > 单股票 ML 策略通常不如多因子规则策略稳定,这是正常现象。")
try:
from reports.storage import save_report
# 组装完整报告
report_lines = []
report_lines.append("# ML 模型训练报告 — {}".format(args.ts_code))
report_lines.append("")
report_lines.append("## 数据概况")
report_lines.append("- 日线: {} 条 ({} ~ {})".format(len(price_df), price_df.index[0], price_df.index[-1]))
report_lines.append("- 因子: {}".format(factor_df.shape[1]))
report_lines.append("- 特征矩阵: {} × {}".format(X.shape[0], X.shape[1]))
report_lines.append("- 标签 (未来{}日收益): mean={:.2f}% std={:.2f}%".format(args.lookahead, y.mean(), y.std()))
report_lines.append("- 训练集: {} 行 | 测试集: {}".format(len(X_train), len(X_test)))
report_lines.append("")
report_lines.append("## LightGBM")
report_lines.append("- 测试集 IC: {:.4f} | 树数: {}".format(ic_lgb, lgb_model.n_estimators_used))
report_lines.append("- 解读: {}".format(
"IC 接近 0,单股票预测信号极弱(正常现象)" if abs(ic_lgb) < 0.03
else "IC {:.3f} 有微弱预测能力".format(ic_lgb) if abs(ic_lgb) < 0.08
else "IC {:.3f} 有显著预测能力".format(ic_lgb)))
report_lines.append("")
report_lines.append("### 特征重要性 (Top 10)")
report_lines.append("| 特征 | 重要性 |")
report_lines.append("|------|--------|")
for _, row in imp.head(10).iterrows():
report_lines.append("| {} | {:.1f}% |".format(row["feature"], row["importance_pct"]))
report_lines.append("")
report_lines.append("## CatBoost")
report_lines.append("- 测试集 IC: {:.4f} | 树数: {}".format(ic_cb, cb_model.n_estimators_used))
report_lines.append("- 解读: {}".format(
"IC 接近 0,两者结论一致:单股票短期收益很难预测" if abs(ic_cb) < 0.03
else "IC {:.3f}".format(ic_cb)))
cb_imp = cb_model.get_feature_importance()
report_lines.append("")
report_lines.append("### 特征重要性 (Top 10)")
report_lines.append("| 特征 | 重要性 |")
report_lines.append("|------|--------|")
for _, row in cb_imp.head(10).iterrows():
report_lines.append("| {} | {:.1f}% |".format(row["feature"], row["importance_pct"]))
report_lines.append("")
report_lines.append("## 回测对比")
report_lines.append("")
# 将 DataFrame 转为 MD 管道表格
if result is not None and not result.empty:
cols = result.columns.tolist()
report_lines.append("| model | " + " | ".join(cols) + " |")
report_lines.append("|" + "|".join(["------"] * (len(cols) + 1)) + "|")
for idx, row in result.iterrows():
vals = []
for c in cols:
v = row[c]
vals.append("{:.2f}".format(v) if isinstance(v, (int, float)) and not np.isnan(v) else str(v) if not (isinstance(v, float) and np.isnan(v)) else "-")
report_lines.append("| " + str(idx) + " | " + " | ".join(vals) + " |")
report_lines.append("")
report_lines.append("> 解读: 正收益+高夏普=模型有效;负收益=需更多特征或换截面预测。单股票 ML 策略通常不如多因子规则策略稳定。")
else:
report_lines.append("无回测数据")
report_lines.append("")
report_lines.append("> 解读: 正收益+高夏普=模型有效;负收益=需更多特征或换截面预测。单股票 ML 策略通常不如多因子规则策略稳定,这是正常现象。")
save_report("\n".join(report_lines), "ML 模型训练报告",
subject_type="stock", subject_code=args.ts_code)
print("\n 报告已存入 DB")
except Exception as e:
print("\n [WARN] 报告入库失败: {}".format(e))
print("\n" + "=" * 60)
print("Sprint 4 验证完成")
print("=" * 60)
if __name__ == "__main__":
main()
+105
View File
@@ -0,0 +1,105 @@
"""
Sprint 3 验证脚本 — Optuna 参数优化。
用法:
python cli/demo_optimizer.py
python cli/demo_optimizer.py --ts_code 600519.SH --trials 100
"""
import sys, os, argparse
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from data.data_manager import DataManager
from factors.engine import FactorEngine
from factors.registry import get_factor
from backtest.vectorbt.engine import VectorBTEngine
from backtest.strategies.rsi_mean_revert import RSIMeanRevertStrategy
from optimizer.engine import OptunaEngine
from optimizer.space import rsi_revert_space
def main():
p = argparse.ArgumentParser(description="Sprint 3 — Optuna 参数优化验证")
p.add_argument("--ts_code", default="000001.SZ", help="测试股票代码")
p.add_argument("--trials", type=int, default=200, help="试验次数(默认: 200")
args = p.parse_args()
print("=" * 60)
print("Sprint 3 — Optuna 参数优化验证")
print("=" * 60)
print("\n[1/4] 初始化...")
dm = DataManager(); dm.init_db()
engine_fe = FactorEngine(dm)
bt_engine = VectorBTEngine()
opt_engine = OptunaEngine(bt_engine)
price_df = dm.get_daily(args.ts_code).set_index("trade_date").sort_index()
factor_df = engine_fe.compute(args.ts_code, [get_factor("rsi_14")])
print(" 数据: {} 条日线, {} 个因子".format(len(price_df), factor_df.shape[1]))
print("[OK]")
print("\n[2/4] RSI 反转策略参数寻优 (sharpe, {} trials)...".format(args.trials))
result = opt_engine.optimize(
strategy_class=RSIMeanRevertStrategy,
search_space=rsi_revert_space,
price_df=price_df, factor_df=factor_df,
metric="sharpe", n_trials=args.trials,
)
print(result.summary())
if result.param_importance:
print(" 参数重要性:")
for k, v in sorted(result.param_importance.items(), key=lambda x: -x[1]):
print(" {}: {:.4f}".format(k, v))
print("\n[3/4] 默认参数 vs 最优参数 对比...")
strategies_compare = [
("默认(30/70)", RSIMeanRevertStrategy(oversold=30, overbought=70)),
("最优", RSIMeanRevertStrategy(**result.best_params)),
]
reports = {}
for name, s in strategies_compare:
reports[name] = bt_engine.run(s, price_df, factor_df)
metrics = [("总收益(%)", "total_return"), ("年化CAGR(%)", "cagr"), ("最大回撤(%)", "max_drawdown"),
("夏普比率", "sharpe_ratio"), ("卡玛比率", "calmar_ratio"), ("胜率(%)", "win_rate"),
("盈利因子", "profit_factor"), ("交易笔数", "total_trades")]
print(" {:<18s} {:>12s} {:>12s}".format("指标", "默认(30/70)", "最优"))
print(" " + "-" * 42)
for label, attr in metrics:
vals = []
for r in reports.values():
v = getattr(r, attr)
vals.append("{:.2f}".format(v) if isinstance(v, float) else str(v))
print(" {:<18s} {:>12s} {:>12s}".format(label, vals[0], vals[1]))
print("\n[4/4] Walk-Forward 滚动窗口验证...")
try:
wf = opt_engine.optimize_walk_forward(
RSIMeanRevertStrategy, rsi_revert_space,
price_df, factor_df, metric="sharpe", n_trials=min(80, args.trials),
train_window=252 * 3, test_window=252,
)
print(wf.summary())
except Exception as e:
print(" [SKIP] Walk-Forward 异常: {}".format(e))
# 存入 DB
try:
from reports.storage import save_report
from datetime import datetime
lines = ["## 参数优化 — {}".format(args.ts_code),
result.summary(), "",
"### 默认 vs 最优对比"]
save_report("\n".join(lines), "参数优化", subject_type="stock", subject_code=args.ts_code)
print("\n 报告已存入 DB")
except Exception as e:
print("\n [WARN] 报告入库失败: {}".format(e))
print("\n" + "=" * 60)
print("Sprint 3 验证完成")
print("=" * 60)
if __name__ == "__main__":
main()
+103
View File
@@ -0,0 +1,103 @@
"""
Sprint 5 验证脚本 — 情绪因子快速验证。
用法:
python cli/demo_sentiment.py
python cli/demo_sentiment.py --ts_code 600519.SH
python cli/demo_sentiment.py --ts_code 000001.SZ --no-qwen
"""
import sys, os, argparse
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import pandas as pd
from data.data_manager import DataManager
from factors.sentiment.news_source import NewsSource, align_news_to_trading_days
from factors.sentiment.qwen_client import QwenClient
from factors.sentiment.sentiment_engine import SentimentEngine
def main():
p = argparse.ArgumentParser(description="Sprint 5 — 情绪因子快速验证")
p.add_argument("--ts_code", default="000001.SZ", help="测试股票代码")
p.add_argument("--no-qwen", action="store_true", help="跳过 Qwen 情绪分析")
args = p.parse_args()
print("=" * 60)
print("Sprint 5 — Qwen 情绪因子验证")
print("=" * 60)
dm = DataManager(); dm.init_db()
print("\n[1/5] 新闻数据源测试...")
news_src = NewsSource(use_mcp=False)
news_df = news_src.fetch(args.ts_code, max_news=10)
if not news_df.empty:
print(" 获取 {} 条新闻".format(len(news_df)))
for _, row in news_df.head(3).iterrows():
print(" [{}] {}... (source={})".format(row["date"], str(row["title"])[:80], row["source"]))
else:
print(" (无新闻数据)")
print("[OK]")
print("\n[2/5] 日期对齐测试...")
from data.data_manager import DataManager as DM
price = dm.get_daily(args.ts_code) if dm.get_daily(args.ts_code) is not None else dm.get_daily("000001.SZ")
if price is not None and not price.empty:
price = price.set_index("trade_date").sort_index()
daily_idx = pd.to_datetime(price.index, format="%Y%m%d", errors="coerce")
test_news = pd.DataFrame({"date": ["20240601", "20240602", "20240603"],
"title": ["周六新闻", "周日新闻", "周一新闻"],
"content": [""] * 3, "source": ["test"] * 3, "url": [""] * 3})
aligned = align_news_to_trading_days(test_news, daily_idx)
for _, row in aligned.iterrows():
print(" {}: {}".format(row["title"], row["date"]))
print("[OK]")
print("\n[3/5] Qwen 客户端状态...")
client = QwenClient()
has_api = bool(client.api_key) or bool(client.local_base_url)
if has_api and not args.no_qwen:
mode = "本地Ollama/{}".format(client.local_model) if client.local_base_url else "DashScope/{}".format(client.model)
print(" 模式: {}".format(mode))
else:
print(" 模式: 未配置或跳过")
print("[OK]")
print("\n[4/5] SentimentEngine 全链路 (max_news=10)...")
try:
sent_engine = SentimentEngine(dm, qwen_client=client, news_source=news_src)
sent_df = sent_engine.compute(args.ts_code, max_news=10)
if sent_df is not None and not sent_df.empty:
valid = sent_df.dropna(how="all")
print(" 情绪因子: {}".format(list(sent_df.columns)))
print(" 有效行: {}/{}".format(len(valid), len(sent_df)))
if not valid.empty:
print(valid.tail(5).round(4).to_string())
except Exception as e:
print(" [WARN] {}".format(e))
print("[OK]")
print("\n[5/5] 分析范围解析...")
scope = sent_engine.get_scope_stocks()
print(" 成分股数量: {}".format(len(scope)))
if scope:
print(" 示例: {}".format(", ".join(scope[:5])))
print("[OK]")
try:
from reports.storage import save_report
save_report("## 情绪因子验证 — {}\n\n- 新闻源: OK\n- 日期对齐: OK\n- Qwen: {}".format(
args.ts_code, "就绪" if has_api else "未配置"),
"情绪因子验证", subject_type="stock", subject_code=args.ts_code)
print(" 报告已存入 DB")
except Exception:
pass
print("\n" + "=" * 60)
print("Sprint 5 验证完成")
print("=" * 60)
if __name__ == "__main__":
main()
+453
View File
@@ -0,0 +1,453 @@
"""
情绪因子详细运行过程演示。
用法:
# 默认:000001.SZ,最近30天,指数范围
python cli/demo_sentiment_detail.py
# 指定股票代码和日期
python cli/demo_sentiment_detail.py --ts_code 600519.SH --date 20260603
python cli/demo_sentiment_detail.py --ts_code 000001.SZ,600519.SH,300750.SZ
# 指定日期范围
python cli/demo_sentiment_detail.py --start 20260501 --end 20260603
# 分析指定指数成分股
python cli/demo_sentiment_detail.py --scope-type index --scope-indexes 000300
# 分析指定板块
python cli/demo_sentiment_detail.py --scope-type sector --scope-sectors 银行,电力设备
# 只使用特定新闻源
python cli/demo_sentiment_detail.py --no-xwlb --no-mcp
python cli/demo_sentiment_detail.py --source akshare
# 跳过 Qwen API 调用(仅演示数据流)
python cli/demo_sentiment_detail.py --no-qwen
"""
import sys, os, json, argparse
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import pandas as pd
import numpy as np
from datetime import datetime
def parse_args():
p = argparse.ArgumentParser(description="情绪因子详细运行过程演示")
p.add_argument("--ts_code", default="000001.SZ",
help="股票代码,多个用逗号分隔(默认: 000001.SZ)")
p.add_argument("--date", default=None,
help="目标日期 YYYYMMDD(默认: 今天)")
p.add_argument("--start", default=None,
help="起始日期 YYYYMMDD(默认: date-30天)")
p.add_argument("--end", default=None,
help="结束日期 YYYYMMDD(默认: date 或今天)")
p.add_argument("--scope-type", default=None,
choices=["index", "sector", "custom", "all"],
help="分析范围类型(覆盖 ts_code")
p.add_argument("--scope-indexes", default="000300",
help="指数代码,逗号分隔(默认: 000300)")
p.add_argument("--scope-sectors", default="",
help="板块名称,逗号分隔")
p.add_argument("--max-news", type=int, default=None,
help="最大新闻条数(默认: .env SENTIMENT_MAX_NEWS_PER_STOCK 或 30")
p.add_argument("--max-analyze", type=int, default=50,
help="Qwen 分析最大条数(默认: 50,控制成本)")
p.add_argument("--no-xwlb", action="store_true", help="禁用新闻联播数据源")
p.add_argument("--no-akshare", action="store_true", help="禁用东方财富数据源")
p.add_argument("--no-mcp", action="store_true", help="禁用 MCP 数据源")
p.add_argument("--source", default=None,
choices=["xwlb", "akshare", "mcp"],
help="仅使用指定数据源")
p.add_argument("--no-qwen", action="store_true", help="跳过 Qwen 分析(仅演示数据流)")
return p.parse_args()
def main():
args = parse_args()
date = args.date or datetime.now().strftime("%Y%m%d")
start = args.start or (datetime.strptime(date, "%Y%m%d") - pd.Timedelta(days=30)).strftime("%Y%m%d")
end = args.end or date
print("=" * 72)
print(" 情绪因子详细运行过程")
print("=" * 72)
print(" 日期: {} ~ {} (目标: {})".format(start, end, date))
# ═══════════════════════════════════════════════════════════════
# Step 0: 初始化
# ═══════════════════════════════════════════════════════════════
print("\n" + "-" * 72)
print("Step 0: 初始化引擎")
print("-" * 72)
from data.data_manager import DataManager
from factors.sentiment.qwen_client import QwenClient
from factors.sentiment.news_source import NewsSource, align_news_to_trading_days
dm = DataManager()
dm.init_db()
client = QwenClient()
has_api = (bool(client.api_key) or bool(client.local_base_url)) and not args.no_qwen
use_xwlb = not args.no_xwlb and (args.source is None or args.source == "xwlb")
use_akshare = not args.no_akshare and (args.source is None or args.source == "akshare")
use_mcp = not args.no_mcp and (args.source is None or args.source == "mcp")
print(" Qwen API: {}".format("DashScope/{}".format(client.model) if (has_api and not client.local_base_url) else (
"本地 Ollama/{}".format(client.local_model) if (has_api and client.local_base_url) else "跳过(--no-qwen 或未配置)")))
print(" 数据源: {}/{}/{}".format(
"xwlb" if use_xwlb else "xwlb(off)",
"akshare" if use_akshare else "akshare(off)",
"mcp" if use_mcp else "mcp(off)",
))
max_news = args.max_news or int(os.getenv("SENTIMENT_MAX_NEWS_PER_STOCK", "30"))
print(" 最大新闻: {} 条 (SENTIMENT_MAX_NEWS_PER_STOCK={})".format(
max_news, os.getenv("SENTIMENT_MAX_NEWS_PER_STOCK", "未设置")))
# 分析范围
if args.scope_type:
from factors.sentiment.sentiment_engine import SentimentEngine
# 临时覆盖环境变量
os.environ["SENTIMENT_SCOPE_TYPE"] = args.scope_type
if args.scope_indexes:
os.environ["SENTIMENT_SCOPE_INDEXES"] = args.scope_indexes
if args.scope_sectors:
os.environ["SENTIMENT_SCOPE_SECTORS"] = args.scope_sectors
sent_tmp = SentimentEngine(dm)
ts_codes = sent_tmp.get_scope_stocks()
print(" 分析范围: {} ({})".format(args.scope_type, len(ts_codes)))
if len(ts_codes) > 10:
print(" 股票示例: {}... (共 {} 只)".format(", ".join(ts_codes[:10]), len(ts_codes)))
else:
print(" 股票: {}".format(", ".join(ts_codes)))
else:
ts_codes = [c.strip() for c in args.ts_code.split(",") if c.strip()]
# ═══════════════════════════════════════════════════════════════
# Step 1: 分别从三个数据源获取新闻
# ═══════════════════════════════════════════════════════════════
print("\n" + "-" * 72)
print("Step 1: 获取新闻 ({} 只股票)".format(len(ts_codes)))
print("-" * 72)
all_raw = []
for ts_code in ts_codes:
print("\n --- {} ---".format(ts_code))
xwlb_raw = pd.DataFrame()
ak_raw = pd.DataFrame()
mcp_raw = pd.DataFrame()
if use_xwlb:
try:
xwlb_src = NewsSource(use_akshare=False, use_mcp=False)
xwlb_raw = xwlb_src.fetch(ts_code, start=start, end=end, max_news=max_news * 3)
print(" 新闻联播(DB xwlb_daily_ext): {} 条 (news_date范围: {}-1~{}-1)".format(
len(xwlb_raw), start, end))
except Exception as e:
print(" 新闻联播: 获取失败 ({})".format(e))
if use_akshare:
try:
ak_src = NewsSource(use_xwlb=False, use_mcp=False)
ak_raw = ak_src.fetch(ts_code, start=start, end=end, max_news=max_news)
print(" 东方财富(AkShare stock_news_em): {}".format(len(ak_raw)))
except Exception as e:
print(" 东方财富: 获取失败 ({})".format(e))
if use_mcp:
try:
mcp_src = NewsSource(use_akshare=False, use_xwlb=False, use_mcp=True)
mcp_raw = mcp_src.fetch(ts_code, start=start, end=end, max_news=max_news)
print(" MCP(trendradar-news): {}".format(len(mcp_raw)))
except Exception as e:
print(" MCP: 获取失败 ({})".format(e))
all_raw.append((ts_code, xwlb_raw, ak_raw, mcp_raw))
# 合并所有股票的结果
frames = []
for _, x, a, m in all_raw:
for df in [x, a, m]:
if not df.empty:
frames.append(df)
raw_news = pd.concat(frames, ignore_index=True) if frames else pd.DataFrame()
if not raw_news.empty:
raw_news = raw_news.drop_duplicates(subset=["title", "date"])
raw_news = raw_news.sort_values("date", ascending=False)
print("\n [汇总] 合并去重后: {} 条新闻".format(len(raw_news)))
if raw_news.empty:
print(" (无新闻数据)")
return
src_counts = raw_news["source"].value_counts()
for src, cnt in src_counts.items():
if src == "xwlb":
label = "新闻联播(DB)"
elif src.startswith("akshare"):
label = "东方财富(AkShare)"
elif src.startswith("mcp"):
label = "MCP(trendradar)"
else:
label = src
print(" {}: {}".format(label, cnt))
# ═══════════════════════════════════════════════════════════════
# Step 2: 新闻详情(按来源分开展示)
# ═══════════════════════════════════════════════════════════════
print("\n" + "-" * 72)
print("Step 2: 新闻详情(按数据源分开展示)")
print("-" * 72)
def show_news(label, df, limit=6):
if df.empty:
print("\n [{}] (无数据)".format(label))
return
print("\n [{}] {}".format(label, len(df)))
for i, (_, row) in enumerate(df.head(limit).iterrows()):
title = str(row["title"])[:80]
content_preview = str(row["content"])[:100].replace("\n", " ")
print("\n [{}/{}] {} | {}".format(i + 1, len(df), row["date"], title))
if content_preview:
print(" 内容: {}...".format(content_preview))
url = row.get("url", "")
if url:
print(" 链接: {}".format(url[:100]))
show_news("新闻联播 (xwlb_daily_ext)", raw_news[raw_news["source"] == "xwlb"], limit=6)
show_news("东方财富 (AkShare stock_news_em)",
raw_news[raw_news["source"].str.startswith("akshare")], limit=6)
show_news("MCP (trendradar-news)",
raw_news[raw_news["source"].str.startswith("mcp")], limit=6)
# ═══════════════════════════════════════════════════════════════
# Step 3: 日期对齐
# ═══════════════════════════════════════════════════════════════
print("\n" + "-" * 72)
print("Step 3: 日期对齐到交易日")
print("-" * 72)
# 交易日历:优先用指定股票 DB 缓存,否则 fallback 到 000001.SZ
first_code = ts_codes[0]
price = _get_trading_calendar(dm, first_code)
if price is None:
print(" {} 无 DB 缓存, fallback 到 000001.SZ".format(first_code))
price = _get_trading_calendar(dm, "000001.SZ")
if price is None:
print(" 无交易日历可用")
return
print("\n 交易日历: {} ~ {} ({} 条)".format(price.index[0], price.index[-1], len(price)))
daily_idx = pd.to_datetime(price.index, format="%Y%m%d", errors="coerce")
aligned_news = align_news_to_trading_days(raw_news, daily_idx)
for label, prefix in [("新闻联播", "xwlb"), ("东方财富", "akshare"), ("MCP", "mcp")]:
df = aligned_news[aligned_news["source"].str.startswith(prefix) if prefix != "xwlb"
else (aligned_news["source"] == "xwlb")]
if df.empty:
continue
dates = sorted(df["date"].unique())
print("\n [{}] {} 条 → {} 个交易日 ({})".format(label, len(df), len(dates),
" +1day偏移" if prefix == "xwlb" else " 直接对齐"))
print(" 日期: {} ~ {}".format(dates[0], dates[-1]))
row = df.iloc[0]
print(" 示例: {} | {}...".format(row["date"], str(row["title"])[:60]))
# ═══════════════════════════════════════════════════════════════
# Step 4: Qwen 情绪分析
# ═══════════════════════════════════════════════════════════════
print("\n" + "-" * 72)
print("Step 4: Qwen 情绪分析")
print("-" * 72)
if not has_api:
print("\n [SKIP] Qwen API 跳过 (--no-qwen 或未配置)")
print(" 使用模拟数据演示因子计算逻辑...")
sentiment_results = _mock_sentiment(aligned_news)
else:
max_analyze = min(len(aligned_news), args.max_analyze)
analyze_news = aligned_news.head(max_analyze)
print("\n 逐条分析 {} 条新闻...".format(max_analyze))
sentiment_results = []
for i, (_, row) in enumerate(analyze_news.iterrows()):
title = str(row["title"])
content = str(row["content"]) if len(str(row["content"])) > 20 else ""
text = "{}\n{}".format(title, content)
result = client.analyze_sentiment(text)
sentiment_results.append({
"date": row["date"],
"title": title,
"sentiment_score": result.get("sentiment_score", 0),
"confidence": result.get("confidence", 0),
"impact_duration": result.get("impact_duration", "short"),
"key_topics": json.dumps(result.get("key_topics", [])),
"source": row.get("source", ""),
})
s = result["sentiment_score"]
icon = "(+)" if s > 0.2 else ("(-)" if s < -0.2 else "(o)")
print(" [{}/{}] {} {:+.1f} c={:.2f} | {}...".format(
i + 1, max_analyze, icon, s,
result["confidence"], title[:60]))
sent_df = pd.DataFrame(sentiment_results)
if not sent_df.empty:
print("\n 情绪分析汇总 ({} 条):".format(len(sent_df)))
print(" 平均情绪: {:+.3f}".format(sent_df["sentiment_score"].mean()))
pos = (sent_df["sentiment_score"] > 0.1).sum()
neu = ((sent_df["sentiment_score"] >= -0.1) & (sent_df["sentiment_score"] <= 0.1)).sum()
neg = (sent_df["sentiment_score"] < -0.1).sum()
print(" 正面(>0.1): {} 中性(-0.1~0.1): {} 负面(<-0.1): {}".format(pos, neu, neg))
if "source" in sent_df.columns:
for src in sent_df["source"].unique():
src_df = sent_df[sent_df["source"] == src]
label = src[:20]
print(" [{}] {} 条, 平均情绪: {:+.3f}".format(label, len(src_df), src_df["sentiment_score"].mean()))
# ═══════════════════════════════════════════════════════════════
# Step 5: 因子计算 + 结果输出
# ═══════════════════════════════════════════════════════════════
print("\n" + "-" * 72)
print("Step 5-6: 因子计算 + 结果输出")
print("-" * 72)
from factors.sentiment.sentiment_factor import (
NewsSentimentFactor,
SentimentConfidenceFactor,
SentimentMomentumFactor,
)
if sent_df.empty:
print(" (无情绪数据)")
return
factors = [
NewsSentimentFactor(window=5, decay=0.3, sentiment_df=sent_df),
SentimentConfidenceFactor(window=5, sentiment_df=sent_df),
SentimentMomentumFactor(period=5, sentiment_df=sent_df),
]
factor_results = {}
for f in factors:
series = f.calculate(price)
factor_results[f.name] = series
stats = series.dropna()
if not stats.empty:
print(" {}: mean={:+.4f} std={:.4f} valid={}/{}".format(
f.name, stats.mean(), stats.std(), len(stats), len(series)))
else:
print(" {}: (全NaN)".format(f.name))
factor_df = pd.DataFrame(factor_results)
valid = factor_df.dropna(how="all")
if valid.empty:
print("\n (无有效因子值)")
return
recent = valid.tail(20)
print("\n === 最近 {} 个交易日情绪因子值 ({}) ===".format(len(recent), first_code))
print(" {:<12s} {:>12s} {:>12s} {:>12s}".format("交易日", "news_sent_5", "news_conf_5", "sent_delta_5"))
print(" {} {} {} {}".format("-" * 12, "-" * 12, "-" * 12, "-" * 12))
for idx, row in recent.iterrows():
ns = "{:+.4f}".format(row["news_sent_5"]) if not pd.isna(row["news_sent_5"]) else " NaN"
nc = "{:+.4f}".format(row["news_conf_5"]) if not pd.isna(row["news_conf_5"]) else " NaN"
sd = "{:+.4f}".format(row["sent_delta_5"]) if not pd.isna(row["sent_delta_5"]) else " NaN"
print(" {:<12s} {:>12s} {:>12s} {:>12s}".format(idx, ns, nc, sd))
latest = valid.iloc[-1]
print("\n === 最新交易日 ({}) ===".format(valid.index[-1]))
print(" news_sent_5 : {:+.4f} (加权情绪, >0偏正面)".format(latest["news_sent_5"]))
print(" news_conf_5 : {:+.4f} (置信度加权)".format(latest["news_conf_5"]))
# 情绪贡献明细
print("\n === 情绪贡献明细 (最近3天) ===")
latest_date = valid.index[-1]
nearby = sent_df[
(sent_df["date"] >= str(int(latest_date) - 3)) &
(sent_df["date"] <= latest_date)
]
if not nearby.empty:
for _, row in nearby.head(30).iterrows():
s = row["sentiment_score"]
impact = "(+)" if s > 0.2 else ("(-)" if s < -0.2 else "(o)")
src = str(row.get("source", ""))
src_s = "xwlb" if src == "xwlb" else ("ak" if src.startswith("akshare") else "mcp")
print(" {} [{:+.1f}] [{}] {}...".format(
impact, s, src_s, str(row["title"])[:70]))
else:
print(" (无最近3天新闻)")
try:
from reports.storage import save_report
first = ts_codes[0] if ts_codes else "unknown"
lines = ["## 情绪因子详细演示", "股票: {}".format(", ".join(ts_codes[:5])),
"数据源: {}条新闻".format(len(raw_news)),
"情绪: news_sent_5={}".format(
latest["news_sent_5"] if "news_sent_5" in latest else "N/A")]
save_report("\n".join(lines), "情绪因子详细演示", subject_type="stock", subject_code=first)
print(" 报告已存入 DB")
except Exception:
pass
print("\n" + "=" * 72)
print(" 情绪因子演示完成")
print("=" * 72)
def _get_trading_calendar(dm, ts_code):
"""获取交易日历:优先 DB 缓存;无缓存则尝试 sync_daily 补齐。"""
try:
from database.dao import get_latest_trade_date
if not get_latest_trade_date(ts_code):
print(" {} 无 DB 缓存,尝试 sync_daily 补齐...".format(ts_code))
try:
n = dm.sync_daily(ts_code)
print(" sync_daily 完成: {}".format(n))
except Exception as e:
print(" sync_daily 失败: {}".format(e))
return None
daily = dm.get_daily(ts_code)
if daily is not None and not daily.empty:
daily = daily.set_index("trade_date").sort_index()
if len(daily) > 0:
return daily
except Exception as e:
print(" 获取交易日历异常: {}".format(e))
return None
def _mock_sentiment(news_df):
results = []
for _, row in news_df.iterrows():
title = str(row["title"]).lower()
pos_words = ["利好", "增长", "突破", "创新高", "盈利", "上升", "支持", "回购", "增持", "分红"]
neg_words = ["利空", "下跌", "亏损", "处罚", "减持", "诉讼", "退市", "警告", "暴跌", "违约"]
pos = sum(1 for w in pos_words if w in title)
neg = sum(1 for w in neg_words if w in title)
if pos > neg:
score = min(0.9, 0.1 + pos * 0.2)
elif neg > pos:
score = max(-0.9, -0.1 - neg * 0.2)
else:
score = np.random.uniform(-0.15, 0.15)
results.append({
"date": row["date"], "title": row["title"],
"sentiment_score": round(score, 1),
"confidence": round(np.random.uniform(0.5, 0.9), 2),
"impact_duration": "short", "key_topics": json.dumps([]),
"source": row.get("source", ""),
})
return results
if __name__ == "__main__":
main()
View File
+46
View File
@@ -0,0 +1,46 @@
"""
全局配置模块。
敏感信息优先从环境变量读取,其次使用默认值(仅开发环境)。
"""
import os
from pathlib import Path
from dotenv import load_dotenv
# 确保加载 finance/.env,无论从哪个目录运行
_env_path = Path(__file__).resolve().parent.parent / ".env"
load_dotenv(_env_path)
# ── MariaDB 连接 ──────────────────────────────────────────
# 通过 SSH 隧道连接:shared/script/autossh.sh
MARIADB_CONFIG = {
"host": os.getenv("MAC_DB_HOST", "127.0.0.1"),
"port": int(os.getenv("MAC_DB_PORT", "13306")),
"user": os.getenv("MAC_DB_USER", "myquant"),
"password": os.getenv("MAC_DB_PASSWORD", "_H(lU1_fF*9baRTp"),
"database": os.getenv("MAC_DB_NAME", "myquant"),
"charset": "utf8mb4",
"pool_size": 5,
"pool_recycle": 600, # 10分钟回收,避免 SSH 隧道半开连接
}
# ── 表名前缀 ──────────────────────────────────────────────
TABLE_PREFIX = "mac_"
# 完整表名
TABLE_STOCK_BASIC = f"{TABLE_PREFIX}stock_basic"
TABLE_STOCK_DAILY = f"{TABLE_PREFIX}stock_daily"
TABLE_STOCK_FINANCIAL = f"{TABLE_PREFIX}stock_financial"
TABLE_REPORT = f"{TABLE_PREFIX}report"
# ── AkShare 配置 ──────────────────────────────────────────
AKSHARE_CONFIG = {
"request_timeout": 30,
"retry_times": 3,
"retry_delay": 2, # 秒
}
# ── 默认参数 ──────────────────────────────────────────────
DEFAULT_START_DATE = "20200101"
DEFAULT_END_DATE = None # None 表示当天
View File
+171
View File
@@ -0,0 +1,171 @@
"""
DataManager — 统一数据管理层。
策略/模型层通过 DataManager 获取数据,不直接访问 AkShare/Tushare 或数据库。
优先从 DB 读取,缺失时依次尝试 AkShare → Tushare 拉取并入库。
"""
import time
import pandas as pd
from config.settings import DEFAULT_START_DATE, DEFAULT_END_DATE
from data.sources.akshare_source import AkShareSource
from data.sources.tushare_source import TushareSource
from database import dao
from database.models import create_all_tables
class DataManager:
"""统一数据管理。双数据源:AkShare(主)+ Tushare(备)。"""
def __init__(self):
self._ak_source: AkShareSource | None = None
self._ts_source: TushareSource | None = None
@property
def ak(self) -> AkShareSource:
if self._ak_source is None:
self._ak_source = AkShareSource()
return self._ak_source
@property
def ts(self) -> TushareSource:
if self._ts_source is None:
self._ts_source = TushareSource()
return self._ts_source
# ── 初始化 ────────────────────────────────────────────
def init_db(self) -> None:
create_all_tables()
# ── 内部:双源 try ────────────────────────────────────
def _try_fetch(self, method_name: str, *args, **kwargs):
"""
依次尝试 AkShare → Tushare 调用同一方法名。
method_name: 'fetch_daily' | 'fetch_stock_list' | 'fetch_financial'
返回: (result_df, source_name) 或 (empty_df, None)
如果 ts_code 是指数代码,自动路由到 fetch_index_daily。
"""
from data.sources.akshare_source import is_index_code
# 指数自动路由
if method_name == "fetch_daily" and args:
ts_code = args[0]
if is_index_code(ts_code):
method_name = "fetch_index_daily"
for label, source in [("Tushare", self.ts), ("AkShare", self.ak)]:
try:
if label == "Tushare" and not source.available:
continue
fn = getattr(source, method_name)
df = fn(*args, **kwargs)
if df is not None and not df.empty:
return df, label
except Exception as e:
print(" [{}] {} 失败: {}".format(label, method_name, e))
return pd.DataFrame(), None
# ── 股票列表 ──────────────────────────────────────────
def get_stock_list(self, force_refresh: bool = False) -> pd.DataFrame:
if not force_refresh:
df = dao.query_stock_list()
if not df.empty:
print("[DataManager] 从 DB 读取股票列表: {}".format(len(df)))
return df
print("[DataManager] 拉取股票列表 (AkShare → Tushare)...")
df, src = self._try_fetch("fetch_stock_list")
if df.empty:
print("[DataManager] 所有数据源均无法获取股票列表")
return pd.DataFrame()
print("[DataManager] 股票列表已入库 ({}): {}".format(src, len(df)))
dao.save_stock_list(df)
time.sleep(2)
return df
# ── 日线数据 ──────────────────────────────────────────
def get_daily(
self,
ts_code: str,
start: str | None = None,
end: str | None = None,
force_refresh: bool = False,
) -> pd.DataFrame:
start = start or DEFAULT_START_DATE
end = end or DEFAULT_END_DATE or time.strftime("%Y%m%d")
if not force_refresh:
df = dao.query_daily(ts_code, start, end)
if not df.empty:
return df
# DB 未命中 → 从数据源拉取
df, src = self._try_fetch("fetch_daily", ts_code, start, end)
if df.empty:
print("[WARN] {} 日线获取失败 (AkShare+Tushare 均不可用)".format(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))
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))
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))
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))
return len(df)
def sync_all_daily(self) -> int:
stock_list = self.get_stock_list()
total = 0
for i, ts_code in enumerate(stock_list.index):
try:
total += self.sync_daily(ts_code)
if (i + 1) % 50 == 0:
print("[DataManager] 进度: {}/{}".format(i + 1, len(stock_list)))
time.sleep(1)
except Exception as e:
print("[WARN] {} 同步失败: {}".format(ts_code, e))
print("[DataManager] 全量同步完成,新增 {}".format(total))
return total
# ── 财务数据 ──────────────────────────────────────────
def get_financial(self, ts_code: str) -> pd.DataFrame:
df = dao.query_financial(ts_code)
if not df.empty:
return df
df, src = self._try_fetch("fetch_financial", ts_code)
if not df.empty:
try:
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))
return df
View File
+216
View File
@@ -0,0 +1,216 @@
"""
AkShare 数据源封装。
统一封装 AkShare 调用,返回标准化的 DataFrame。
所有网络请求都在这一层处理,包含重试和容错。
支持个股和指数两种数据接口。
"""
# 指数代码识别:.SH 后缀为主板指数,399xxx 为深证指数
# 注意区分:000001.SZ 是平安银行(个股)000001.SH 是上证指数
def is_index_code(ts_code: str) -> bool:
"""判断是否为指数代码。排除 000xxx.SZ 个股。"""
if not ts_code:
return False
code = ts_code.upper()
# .SH 开头 000 是指数
if code.endswith(".SH") and (code.startswith("000") or code.startswith("399")):
return True
# 深交所 399xxx 指数
if code.startswith("399"):
return True
# 无后缀的纯数字 000xxx(上证指数常见写法)
if code == "000001":
return True
return False
import time
import pandas as pd
import akshare as ak
from config.settings import AKSHARE_CONFIG
class AkShareSource:
"""AkShare 数据源。"""
def __init__(self):
self._timeout = AKSHARE_CONFIG["request_timeout"]
self._retry = AKSHARE_CONFIG["retry_times"]
self._delay = AKSHARE_CONFIG["retry_delay"]
def _retry_call(self, fn, name: str, **kwargs):
"""带重试的 API 调用包装。重试间隔递增。"""
last_err = None
for i in range(self._retry):
try:
return fn(**kwargs)
except Exception as e:
last_err = e
wait = self._delay * (i + 1)
print(f" [retry] {name} 失败 ({e}){wait}s 后重试 ({i + 1}/{self._retry})...")
if i < self._retry - 1:
time.sleep(wait)
raise last_err # type: ignore
# ── 股票列表 ──────────────────────────────────────────
def fetch_stock_list(self) -> pd.DataFrame:
"""获取 A 股股票列表。"""
df = self._retry_call(ak.stock_info_a_code_name, "stock_list")
df = df.rename(columns={
"code": "ts_code",
"name": "name",
})
return df[["ts_code", "name"]]
# ── 日线数据 ──────────────────────────────────────────
def fetch_daily(
self, ts_code: str, start: str, end: str | None = None
) -> pd.DataFrame:
"""
获取单只股票日线行情。
参数:
ts_code: 股票代码,如 '000001'(纯数字格式,AkShare 要求)
start: 起始日期 'YYYYMMDD'
end: 结束日期 'YYYYMMDD'None 表示今天
"""
symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
end = end or time.strftime("%Y%m%d")
df = self._retry_call(
ak.stock_zh_a_hist,
"daily",
symbol=symbol,
period="daily",
start_date=start,
end_date=end,
adjust="qfq", # 前复权
)
if df.empty:
return df
df = df.rename(columns={
"日期": "trade_date",
"开盘": "open",
"收盘": "close",
"最高": "high",
"最低": "low",
"成交量": "vol",
"成交额": "amount",
"振幅": "amplitude",
"涨跌幅": "pct_chg",
"涨跌额": "change",
"换手率": "turnover_rate",
})
df["ts_code"] = ts_code
# AkShare 返回 'YYYY-MM-DD',统一转为 'YYYYMMDD'
df["trade_date"] = df["trade_date"].astype(str).str.replace("-", "")
return df
# ── 指数日线 ──────────────────────────────────────────
def fetch_index_daily(
self, ts_code: str, start: str, end: str | None = None
) -> pd.DataFrame:
"""
获取指数日线行情。
AkShare index_zh_a_hist 接口,symbol 为纯数字(如 '000001')。
"""
symbol = ts_code.replace(".SH", "").replace(".SZ", "").replace(".BJ", "")
end = end or time.strftime("%Y%m%d")
df = self._retry_call(
ak.index_zh_a_hist,
"index_daily",
symbol=symbol, period="daily",
start_date=start, end_date=end,
)
if df.empty:
return df
df = df.rename(columns={
"日期": "trade_date",
"开盘": "open",
"收盘": "close",
"最高": "high",
"最低": "low",
"成交量": "vol",
"成交额": "amount",
"涨跌幅": "pct_chg",
"涨跌额": "change",
})
df["ts_code"] = ts_code
df["trade_date"] = df["trade_date"].astype(str)
return df
# ── 财务数据 ──────────────────────────────────────────
def fetch_financial(self, ts_code: str) -> pd.DataFrame:
"""获取单只股票核心财务指标(同花顺接口)。"""
symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
try:
df = ak.stock_financial_abstract_ths(symbol=symbol)
if df.empty:
return pd.DataFrame()
df = df.rename(columns={
"报告期": "end_date",
"净利润": "net_profit",
"净利润同比增长率": "net_profit_yoy",
"扣非净利润": "deducted_net_profit",
"扣非净利润同比增长率": "deducted_net_profit_yoy",
"营业总收入": "total_revenue",
"营业总收入同比增长率": "total_revenue_yoy",
"基本每股收益": "eps",
"每股净资产": "bvps",
"每股资本公积金": "capital_reserve_ps",
"每股未分配利润": "undistributed_profit_ps",
"每股经营现金流": "ocf_ps",
"销售净利率": "net_profit_margin",
"净资产收益率": "roe",
"净资产收益率-摊薄": "roe_diluted",
"营业周期": "operating_cycle",
"应收账款周转天数": "receivables_days",
"流动比率": "current_ratio",
"速动比率": "quick_ratio",
"保守速动比率": "conservative_quick_ratio",
"产权比率": "equity_ratio",
"资产负债率": "debt_to_assets",
})
df["ts_code"] = ts_code
# 日期格式统一
df["end_date"] = df["end_date"].astype(str).str.replace("-", "")
# 数值列清洗:去掉 万/亿/% 等单位
for col in df.columns:
if col in ("ts_code", "end_date"):
continue
df[col] = self._parse_financial_value(df[col])
return df
except Exception:
return pd.DataFrame()
@staticmethod
def _parse_financial_value(series: pd.Series) -> pd.Series:
"""解析财务数值字符串 '4302.00万', '64.75%', 'False' → float"""
def _parse(v):
if v is None or v == "False" or v == "":
return None
if isinstance(v, (int, float)):
return float(v)
s = str(v).strip()
if not s:
return None
try:
if s.endswith("%"):
return float(s[:-1])
if "" in s:
return float(s.replace("", "")) * 1e4
if "亿" in s:
return float(s.replace("亿", "")) * 1e8
return float(s)
except ValueError:
return None
return series.apply(_parse).astype("float64")
+139
View File
@@ -0,0 +1,139 @@
"""
Tushare 数据源封装。
与 AkShareSource 保持相同接口,作为平替/备份数据源。
"""
import os
import time
import pandas as pd
import tushare as ts
# 确保 .env 已加载(无论从哪个路径导入)
import config.settings # noqa: F401
class TushareSource:
"""Tushare 数据源。"""
def __init__(self, token: str | None = None):
token = token or os.getenv("TUSHARE_TOKEN", "")
if not token:
self._pro = None
else:
ts.set_token(token)
self._pro = ts.pro_api()
@property
def available(self) -> bool:
return self._pro is not None
def _ensure_pro(self):
if self._pro is None:
raise RuntimeError("Tushare token 未配置,请在 .env 中设置 TUSHARE_TOKEN")
# ── 股票列表 ──────────────────────────────────────────
def fetch_stock_list(self) -> pd.DataFrame:
"""获取 A 股股票列表。"""
self._ensure_pro()
fields = "ts_code,name,area,industry,market,list_date,is_hs"
df = self._pro.stock_basic(
exchange="", list_status="L",
fields=fields,
)
if df is None or df.empty:
return pd.DataFrame()
return df[["ts_code", "name"]]
# ── 日线数据 ──────────────────────────────────────────
def fetch_daily(
self, ts_code: str, start: str, end: str | None = None
) -> pd.DataFrame:
"""
获取单只股票日线行情。
参数:
ts_code: 如 '000001.SZ'
start: 'YYYYMMDD'
end: 'YYYYMMDD'
"""
self._ensure_pro()
end = end or time.strftime("%Y%m%d")
df = self._pro.daily(
ts_code=ts_code,
start_date=start,
end_date=end,
fields="ts_code,trade_date,open,high,low,close,pre_close,change,pct_chg,vol,amount,turnover_rate",
)
if df is None or df.empty:
return pd.DataFrame()
# 前复权因子
try:
adj = self._pro.adj_factor(ts_code=ts_code, start_date=start, end_date=end)
if adj is not None and not adj.empty:
df = df.merge(adj[["trade_date", "adj_factor"]], on="trade_date", how="left")
for col in ["open", "high", "low", "close", "pre_close"]:
if col in df.columns and "adj_factor" in df.columns:
df[col] = (df[col] * df["adj_factor"]).round(2)
df = df.drop(columns=["adj_factor"])
except Exception:
pass
df["trade_date"] = df["trade_date"].astype(str)
return df
# ── 指数日线 ──────────────────────────────────────────
def fetch_index_daily(
self, ts_code: str, start: str, end: str | None = None
) -> pd.DataFrame:
"""
获取指数日线行情。
Tushare index_daily 接口。
"""
self._ensure_pro()
end = end or time.strftime("%Y%m%d")
df = self._pro.index_daily(
ts_code=ts_code,
start_date=start, end_date=end,
fields="ts_code,trade_date,open,high,low,close,pre_close,change,pct_chg,vol,amount",
)
if df is None or df.empty:
return pd.DataFrame()
df["trade_date"] = df["trade_date"].astype(str)
return df
# ── 财务数据 ──────────────────────────────────────────
def fetch_financial(self, ts_code: str) -> pd.DataFrame:
"""获取单只股票核心财务指标。"""
self._ensure_pro()
fields = (
"ts_code,end_date,roe,roa,grossprofit_margin,netprofit_margin,"
"debt_to_assets,eps,dt_eps,bps,pe,pb,"
"total_revenue,revenue_yoy,n_income,n_income_yoy"
)
df = self._pro.fina_indicator(
ts_code=ts_code,
fields=fields,
)
if df is None or df.empty:
return pd.DataFrame()
df = df.rename(columns={
"grossprofit_margin": "gross_profit_margin",
"netprofit_margin": "net_profit_margin",
"dt_eps": "eps_diluted",
"bps": "bvps",
"revenue_yoy": "total_revenue_yoy",
"n_income": "net_profit",
"n_income_yoy": "net_profit_yoy",
})
df["end_date"] = df["end_date"].astype(str)
return df
View File
+118
View File
@@ -0,0 +1,118 @@
"""
MariaDB 连接管理。
使用 SQLAlchemy 连接池,通过 SSH 隧道访问远程数据库。
连接失败时自动尝试重建 SSH 隧道并重连。
"""
import os
import subprocess
import time
from sqlalchemy import create_engine, text
from sqlalchemy.orm import sessionmaker
from config.settings import MARIADB_CONFIG
_engine = None
_SessionLocal = None
def _build_url() -> str:
cfg = MARIADB_CONFIG
return (
"mysql+pymysql://{}:{}"
"@{}:{}/{}"
"?charset={}"
).format(cfg["user"], cfg["password"], cfg["host"], cfg["port"], cfg["database"], cfg["charset"])
def _reconnect_ssh() -> bool:
"""自动运行 autossh.sh 重建 SSH 隧道。每次调用都会尝试。"""
candidates = [
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(
os.path.abspath(__file__)))), "shared", "script", "autossh.sh"),
os.path.expanduser("~/Downloads/cc-cursor/shared/script/autossh.sh"),
]
script = None
for c in candidates:
if os.path.exists(c):
script = c
break
if script is None:
print("[DB] SSH 脚本未找到")
return False
try:
print("[DB] 尝试重建 SSH 隧道: {}".format(script))
result = subprocess.run(["bash", script], capture_output=True, text=True, timeout=15)
time.sleep(3) # 等待隧道建立
if result.returncode == 0:
print("[DB] SSH 隧道重建完成")
return True
else:
print("[DB] SSH 隧道重建失败: {}".format(result.stderr[:200]))
return False
except Exception as e:
print("[DB] SSH 执行异常: {}".format(e))
return False
def _test_engine(engine) -> bool:
"""测试引擎连接是否存活。"""
try:
with engine.connect() as conn:
conn.execute(text("SELECT 1"))
return True
except Exception:
return False
def get_engine():
"""获取 SQLAlchemy Engine(懒初始化,断连时自动重连)。"""
global _engine, _SessionLocal
# 已有引擎但连接断开 → dispose 旧池 → SSH 重连 → 重建引擎
if _engine is not None and not _test_engine(_engine):
print("[DB] 连接已断开,尝试自动恢复...")
_engine.dispose() # 杀死所有旧连接
_engine = None
if _reconnect_ssh():
print("[DB] 连接已恢复")
else:
print("[DB] 自动恢复失败,将无法连接数据库")
# 创建引擎(首次或重建后)
if _engine is None:
cfg = MARIADB_CONFIG
_engine = create_engine(
_build_url(),
pool_size=cfg["pool_size"],
pool_recycle=cfg["pool_recycle"],
pool_pre_ping=True, # 每次取连接前先 SELECT 1 验证存活
echo=False,
)
# Session factory 下次获取时自动绑定新引擎
_SessionLocal = None
return _engine
def get_session():
"""获取一个新的数据库会话。"""
global _SessionLocal
if _SessionLocal is None:
_SessionLocal = sessionmaker(bind=get_engine())
return _SessionLocal()
def test_connection() -> bool:
"""测试数据库连接是否正常。"""
try:
engine = get_engine()
return _test_engine(engine)
except Exception as e:
print("[ERROR] 数据库连接失败: {}".format(e))
return False
+133
View File
@@ -0,0 +1,133 @@
"""
数据访问对象。
提供 DataFrame 级别的读写操作,屏蔽底层 ORM/SQL 细节。
"""
import pandas as pd
from sqlalchemy import text
from database.connection import get_engine
from database.models import StockBasic, StockDaily, StockFinancial, Report
# DB 表列名,供 DataManager 在写入前筛选
_DAILY_COLS = [
"ts_code", "trade_date", "open", "high", "low", "close",
"pre_close", "change", "pct_chg", "vol", "amount", "turnover_rate",
]
_FINA_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",
]
def _df_to_db(df: pd.DataFrame, model_class, replace: bool = False) -> int:
"""将 DataFrame 写入对应表,返回写入行数。"""
if df.empty:
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
# ── StockBasic ─────────────────────────────────────────────
def save_stock_list(df: pd.DataFrame) -> int:
"""保存股票列表(replace 模式)。"""
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)
def query_stock_list() -> pd.DataFrame:
"""查询全部股票列表。"""
engine = get_engine()
return pd.read_sql(f"SELECT * FROM {StockBasic.__tablename__}", con=engine).set_index("ts_code")
# ── StockDaily ─────────────────────────────────────────────
def save_daily(df: pd.DataFrame) -> int:
"""批量写入日线数据。先删旧再插新,避免主键冲突。"""
cols = [
"ts_code", "trade_date", "open", "high", "low", "close",
"pre_close", "change", "pct_chg", "vol", "amount", "turnover_rate",
]
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)
def query_daily(ts_code: str, start: str | None = None, end: str | None = None) -> pd.DataFrame:
"""按股票代码和日期范围查询日线。"""
engine = get_engine()
table = StockDaily.__tablename__
sql = f"SELECT * FROM {table} WHERE ts_code = :ts_code"
params = {"ts_code": ts_code}
if start:
sql += " AND trade_date >= :start"
params["start"] = start
if end:
sql += " AND trade_date <= :end"
params["end"] = end
sql += " ORDER BY trade_date ASC"
df = pd.read_sql(text(sql), con=engine, params=params)
if not df.empty:
df["trade_date"] = df["trade_date"].astype(str)
return df
def get_latest_trade_date(ts_code: str) -> str | None:
"""获取某股票在数据库中的最新交易日。"""
engine = get_engine()
table = StockDaily.__tablename__
sql = f"SELECT MAX(trade_date) FROM {table} WHERE ts_code = :ts_code"
with engine.connect() as conn:
result = conn.execute(text(sql), {"ts_code": ts_code}).scalar()
return result
# ── StockFinancial ─────────────────────────────────────────
def save_financial(df: pd.DataFrame) -> int:
"""批量写入财务数据(replace 模式:同报告期覆盖更新)。"""
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)
def query_financial(ts_code: str) -> pd.DataFrame:
"""查询某股票全部财务数据。"""
engine = get_engine()
table = StockFinancial.__tablename__
sql = f"SELECT * FROM {table} WHERE ts_code = :ts_code ORDER BY end_date DESC"
return pd.read_sql(text(sql), con=engine, params={"ts_code": ts_code})
+100
View File
@@ -0,0 +1,100 @@
"""
ORM 模型定义。
所有表使用 mac_ 前缀,与现有表隔离。
"""
from sqlalchemy import Column, String, Date, DateTime, Float, BigInteger, Index, PrimaryKeyConstraint, Text
from sqlalchemy.orm import DeclarativeBase
from config.settings import TABLE_STOCK_BASIC, TABLE_STOCK_DAILY, TABLE_STOCK_FINANCIAL, TABLE_REPORT
class Base(DeclarativeBase):
pass
class StockBasic(Base):
"""股票基本信息表。"""
__tablename__ = TABLE_STOCK_BASIC
ts_code = Column(String(16), primary_key=True, comment="股票代码(如 000001.SZ")
name = Column(String(32), comment="股票名称")
area = Column(String(16), comment="地区")
industry = Column(String(32), comment="行业")
market = Column(String(8), comment="市场(主板/创业板/科创板)")
list_date = Column(String(8), comment="上市日期")
is_hs = Column(String(1), comment="是否沪深港通")
class StockDaily(Base):
"""日线行情表。"""
__tablename__ = TABLE_STOCK_DAILY
__table_args__ = (
PrimaryKeyConstraint("ts_code", "trade_date"),
Index("idx_mac_daily_ts_code", "ts_code"),
Index("idx_mac_daily_trade_date", "trade_date"),
)
ts_code = Column(String(16), comment="股票代码")
trade_date = Column(String(8), comment="交易日期")
open = Column(Float, comment="开盘价")
high = Column(Float, comment="最高价")
low = Column(Float, comment="最低价")
close = Column(Float, comment="收盘价")
pre_close = Column(Float, comment="昨收价")
change = Column(Float, comment="涨跌额")
pct_chg = Column(Float, comment="涨跌幅(%)")
vol = Column(Float, comment="成交量(手)")
amount = Column(Float, comment="成交额(千元)")
turnover_rate = Column(Float, comment="换手率(%)")
class StockFinancial(Base):
"""财务数据表(同花顺核心指标)。"""
__tablename__ = TABLE_STOCK_FINANCIAL
__table_args__ = (
PrimaryKeyConstraint("ts_code", "end_date"),
Index("idx_mac_fina_ts_code", "ts_code"),
)
ts_code = Column(String(16), comment="股票代码")
end_date = Column(String(8), comment="报告期 YYYYMMDD")
eps = Column(Float, comment="基本每股收益")
bvps = Column(Float, comment="每股净资产")
roe = Column(Float, comment="净资产收益率(%)")
roe_diluted = Column(Float, comment="净资产收益率-摊薄(%)")
net_profit_margin = Column(Float, comment="销售净利率(%)")
debt_to_assets = Column(Float, comment="资产负债率(%)")
current_ratio = Column(Float, comment="流动比率")
quick_ratio = Column(Float, comment="速动比率")
total_revenue = Column(Float, comment="营业总收入")
total_revenue_yoy = Column(Float, comment="营业总收入同比增长率(%)")
net_profit = Column(Float, comment="净利润")
net_profit_yoy = Column(Float, comment="净利润同比增长率(%)")
class Report(Base):
"""报告存储表。"""
__tablename__ = TABLE_REPORT
__table_args__ = (
Index("idx_mac_report_date", "report_date"),
Index("idx_mac_report_subject", "subject_type", "subject_code"),
Index("idx_mac_report_active", "is_active"),
)
id = Column(BigInteger, primary_key=True, autoincrement=True, comment="主键")
report_date = Column(Date, nullable=False, comment="报告日期")
title = Column(String(256), nullable=False, comment="报告标题")
subject_type = Column(String(32), comment="研究对象类型: stock/index/sector/portfolio/daily")
subject_code = Column(String(64), comment="研究对象代码")
content = Column(Text, comment="报告内容 (markdown)")
created_at = Column(DateTime, comment="报告生成时间")
is_active = Column(Float, default=1.0, comment="1=有效, 0=已失效")
def create_all_tables():
"""创建所有 mac_ 开头的表。"""
engine = __import__("database.connection", fromlist=["get_engine"]).get_engine()
Base.metadata.create_all(engine)
print("[OK] 所有表创建完成")
View File
View File
+42
View File
@@ -0,0 +1,42 @@
"""
因子抽象基类。
所有因子必须继承 BaseFactor,实现 calculate(df) → pd.Series。
"""
from abc import ABC, abstractmethod
import pandas as pd
class BaseFactor(ABC):
"""因子抽象基类。
属性:
name: 因子名称,如 'momentum_20'
category: 'technical' | 'fundamental' | 'sentiment'
required_columns: 计算所需的 DataFrame 列名列表
"""
name: str = ""
category: str = ""
@abstractmethod
def calculate(self, df: pd.DataFrame) -> pd.Series:
"""计算因子值。
参数:
df: 日线 DataFrameindex 为 trade_date
列至少包含 OHLCV 等基础字段。
返回:
pd.Seriesindex 与 df 对齐,值为因子值。
"""
...
def get_required_columns(self) -> list[str]:
"""返回计算所需的列名列表。子类可覆盖。"""
return ["close"]
def __repr__(self) -> str:
return f"{self.__class__.__name__}(name='{self.name}')"
+192
View File
@@ -0,0 +1,192 @@
"""
FactorEngine — 因子计算引擎。
批量计算因子,处理技术/基本面/情绪因子的不同数据需求。
"""
import copy
import pandas as pd
from factors.base import BaseFactor
from factors.fundamental.roe import ROEFactor
from factors.fundamental.pe_pb import PEFactor, PBFactor, EPFactor
FUNDAMENTAL_FACTOR_TYPES = (ROEFactor, PEFactor, PBFactor, EPFactor)
def _is_sentiment(factor: BaseFactor) -> bool:
return getattr(factor, "category", "") == "sentiment"
class FactorEngine:
"""因子计算引擎。"""
def __init__(self, data_manager, sentiment_engine=None):
"""
参数:
data_manager: DataManager 实例。
sentiment_engine: SentimentEngine 实例(可选,启用情绪因子时需提供)。
"""
self._dm = data_manager
self._sentiment_engine = sentiment_engine
self._financial_cache: dict[str, pd.DataFrame] = {}
def _get_financial(self, ts_code: str) -> pd.DataFrame:
"""获取财务数据(带缓存)。"""
if ts_code not in self._financial_cache:
df = self._dm.get_financial(ts_code)
self._financial_cache[ts_code] = df
return self._financial_cache[ts_code]
def _resolve_factors(
self, factors: list[BaseFactor], fina: pd.DataFrame
) -> list[BaseFactor]:
"""为每个股票 clone 基本面因子并注入财务数据。"""
resolved = []
for f in factors:
if isinstance(f, FUNDAMENTAL_FACTOR_TYPES):
f = copy.copy(f)
f._financial_df = fina
resolved.append(f)
return resolved
def compute(
self,
ts_code: str,
factors: list[BaseFactor],
) -> pd.DataFrame:
"""
对单只股票计算多个因子。
参数:
ts_code: 如 '000001.SZ'
factors: 因子实例列表
返回:
DataFrameindex=trade_datecolumns=因子名
"""
if not factors:
return pd.DataFrame()
# 分离情绪因子(通过 SentimentEngine 处理)
sent_factors = [f for f in factors if _is_sentiment(f)]
other_factors = [f for f in factors if not _is_sentiment(f)]
# 收集所有需要的列
required_cols = set()
has_fundamental = False
for f in other_factors:
required_cols.update(f.get_required_columns())
if isinstance(f, FUNDAMENTAL_FACTOR_TYPES):
has_fundamental = True
# 获取日线数据
daily = self._dm.get_daily(ts_code)
if daily.empty:
return pd.DataFrame()
daily = daily.set_index("trade_date").sort_index()
# 获取财务数据(如有基本面因子)
fina = self._dm.get_financial(ts_code) if has_fundamental else pd.DataFrame()
# 为当前股票解析因子(clone 基本面因子注入财务数据)
resolved_factors = self._resolve_factors(other_factors, fina)
# 逐因子计算
results = {}
for factor in resolved_factors:
try:
series = factor.calculate(daily)
results[factor.name] = series.astype("float64")
except Exception as e:
print(f"[WARN] 因子 {factor.name} 计算失败 ({ts_code}): {e}")
results[factor.name] = pd.Series(float("nan"), index=daily.index)
# 情绪因子:通过 SentimentEngine 计算后合并
if sent_factors and self._sentiment_engine:
try:
sent_df = self._sentiment_engine.compute(ts_code)
for f in sent_factors:
if f.name in sent_df.columns:
results[f.name] = sent_df[f.name]
else:
results[f.name] = pd.Series(float("nan"), index=daily.index)
except Exception as e:
print(f"[WARN] 情绪因子计算失败 ({ts_code}): {e}")
for f in sent_factors:
results[f.name] = pd.Series(float("nan"), index=daily.index)
factor_df = pd.DataFrame(results)
factor_df.index.name = "trade_date"
return factor_df
def compute_batch(
self,
ts_codes: list[str],
factors: list[BaseFactor],
) -> dict[str, pd.DataFrame]:
"""
批量计算多只股票的因子。
返回:
{ts_code: factor_df}
"""
results = {}
total = len(ts_codes)
for i, ts_code in enumerate(ts_codes):
try:
results[ts_code] = self.compute(ts_code, factors)
except Exception as e:
print(f"[WARN] {ts_code} 因子计算失败: {e}")
results[ts_code] = pd.DataFrame()
if (i + 1) % 50 == 0:
print(f"[FactorEngine] 进度: {i + 1}/{total}")
return results
def compute_universe(
self,
factors: list[BaseFactor],
date: str,
ts_codes: list[str] | None = None,
) -> pd.DataFrame:
"""
计算全市场某一天的因子截面。
参数:
factors: 因子列表
date: 目标日期 'YYYYMMDD'
ts_codes: 股票列表,None 表示全部
返回:
DataFrameindex=ts_codecolumns=因子名
"""
if ts_codes is None:
stocks = self._dm.get_stock_list()
ts_codes = list(stocks.index)
rows = []
for ts_code in ts_codes:
daily = self._dm.get_daily(ts_code)
if daily.empty:
continue
daily = daily.set_index("trade_date")
if date not in daily.index:
continue
row = {"ts_code": ts_code}
fina = self._get_financial(ts_code)
resolved = self._resolve_factors(factors, fina)
for factor in resolved:
try:
series = factor.calculate(daily)
row[factor.name] = series.get(date, float("nan"))
except Exception:
row[factor.name] = float("nan")
rows.append(row)
if not rows:
return pd.DataFrame()
result = pd.DataFrame(rows).set_index("ts_code")
return result
+113
View File
@@ -0,0 +1,113 @@
"""
PE / PB 估值因子
基于日线收盘价 + 财务数据EPS/每股净资产计算
"""
import pandas as pd
from factors.base import BaseFactor
class PEFactor(BaseFactor):
"""
市盈率因子 = close / eps
eps 来自财务数据中的 'eps' 列或 TTM EPS
因子值越大表示估值越贵
"""
category = "fundamental"
def __init__(self, financial_df: pd.DataFrame | None = None):
"""
参数:
financial_df: 'end_date' 'eps' DataFrame
"""
self._financial_df = financial_df
self.name = "pe"
def calculate(self, df: pd.DataFrame) -> pd.Series:
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")
close = df["close"]
return close / eps_series.replace(0, float("nan"))
def get_required_columns(self) -> list[str]:
return ["close"]
class PBFactor(BaseFactor):
"""
市净率因子 = close / bvps每股净资产
因子值越大表示估值越贵
"""
category = "fundamental"
def __init__(self, financial_df: pd.DataFrame | None = None):
self._financial_df = financial_df
self.name = "pb"
def calculate(self, df: pd.DataFrame) -> pd.Series:
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")
return df["close"] / bvps_series.replace(0, float("nan"))
def get_required_columns(self) -> list[str]:
return ["close"]
class EPFactor(BaseFactor):
"""
盈利收益率因子 = eps / close = 1 / PE
值越大表示估值越便宜适合与动量等因子同向排序
"""
category = "fundamental"
def __init__(self, financial_df: pd.DataFrame | None = None):
self._financial_df = financial_df
self.name = "ep"
def calculate(self, df: pd.DataFrame) -> pd.Series:
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")
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")
+104
View File
@@ -0,0 +1,104 @@
"""
ROE 因子
"""
import pandas as pd
from factors.base import BaseFactor
class ROEFactor(BaseFactor):
"""
ROE 因子
从财务数据提取 ROE 并映射到日线
需要 df 中包含 'roe' FactorEngine 合并财务数据后传入
或将 financial_df 直接传入构造函数
"""
category = "fundamental"
def __init__(self, financial_df: pd.DataFrame | None = None):
"""
参数:
financial_df: 财务数据 DataFramecolumns 'end_date', 'roe'
None 时需在 df 参数中直接提供 roe
"""
self._financial_df = financial_df
self.name = "roe"
def calculate(self, df: pd.DataFrame) -> pd.Series:
if "roe" in df.columns:
return df["roe"].copy()
if self._financial_df is None or self._financial_df.empty:
return pd.Series(float("nan"), index=df.index)
return self._map_financial_to_daily(
df, self._financial_df, "roe"
)
@staticmethod
def _map_financial_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]
else:
pass # 最新一期覆盖所有后续日期
result[mask] = fina[column].iloc[i]
return result.astype("float64")
class ROETTMDeltaFactor(BaseFactor):
"""ROE 同比变化(当前 ROE - 去年同期 ROE)。"""
category = "fundamental"
def __init__(self, financial_df: pd.DataFrame | None = None):
self._financial_df = financial_df
self.name = "roe_delta"
def calculate(self, df: pd.DataFrame) -> pd.Series:
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")
# 按年分组计算 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:
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
roe_delta[mask] = delta
return roe_delta.astype("float64")
+127
View File
@@ -0,0 +1,127 @@
"""
因子注册表
通过名称获取因子实例方便回测和策略配置时引用
"""
from factors.base import BaseFactor
from factors.technical.momentum import MomentumFactor
from factors.technical.rsi import RSIFactor
from factors.technical.macd import MACDFactor
from factors.technical.volume import VolumeFactor, VolumeChangeFactor
from factors.technical.bollinger import BollingerFactor, BollingerWidthFactor
from factors.technical.atr import ATRFactor, ATRRatioFactor
from factors.technical.ma_cross import MACrossFactor, MADeviationFactor
from factors.technical.volatility import VolatilityFactor, DownsideVolatilityFactor
from factors.technical.turnover import TurnoverFactor, TurnoverChangeFactor
from factors.technical.amplitude import AmplitudeFactor
from factors.fundamental.roe import ROEFactor
from factors.fundamental.pe_pb import PEFactor, PBFactor, EPFactor
from factors.sentiment.sentiment_factor import (
NewsSentimentFactor,
SentimentMomentumFactor,
SentimentConfidenceFactor,
)
# ── 内置因子工厂函数 ──────────────────────────────────────
_BUILTIN_FACTORIES: dict[str, callable] = { # type: ignore
# 动量
"momentum_5": lambda: MomentumFactor(period=5),
"momentum_10": lambda: MomentumFactor(period=10),
"momentum_20": lambda: MomentumFactor(period=20),
"momentum_60": lambda: MomentumFactor(period=60),
# RSI
"rsi_7": lambda: RSIFactor(period=7),
"rsi_14": lambda: RSIFactor(period=14),
# MACD
"macd": lambda: MACDFactor(),
"macd_5_35_5": lambda: MACDFactor(fast=5, slow=35, signal=5),
# 量价
"vol_ratio_5": lambda: VolumeFactor(period=5),
"vol_ratio_20": lambda: VolumeFactor(period=20),
"vol_chg_5": lambda: VolumeChangeFactor(period=5),
# 布林
"boll": lambda: BollingerFactor(),
"boll_width": lambda: BollingerWidthFactor(),
# ATR
"atr_14": lambda: ATRFactor(period=14),
"atr_ratio_14": lambda: ATRRatioFactor(period=14),
# 均线
"ma_cross_5_20": lambda: MACrossFactor(fast=5, slow=20),
"ma_cross_10_60": lambda: MACrossFactor(fast=10, slow=60),
"ma_dev_20": lambda: MADeviationFactor(period=20),
"ma_dev_60": lambda: MADeviationFactor(period=60),
# 波动率
"volatility_20": lambda: VolatilityFactor(period=20),
"volatility_60": lambda: VolatilityFactor(period=60),
"down_vol_20": lambda: DownsideVolatilityFactor(period=20),
# 换手率
"turnover_5": lambda: TurnoverFactor(period=5),
"turnover_chg_5": lambda: TurnoverChangeFactor(period=5),
# 振幅
"amplitude_5": lambda: AmplitudeFactor(period=5),
"amplitude_20": lambda: AmplitudeFactor(period=20),
# 基本面
"roe": lambda: ROEFactor(),
"pe": lambda: PEFactor(),
"pb": lambda: PBFactor(),
"ep": lambda: EPFactor(),
# 情绪
"news_sent_5": lambda: NewsSentimentFactor(window=5),
"news_sent_20": lambda: NewsSentimentFactor(window=20),
"news_conf_5": lambda: SentimentConfidenceFactor(window=5),
"sent_delta_5": lambda: SentimentMomentumFactor(period=5),
}
# ── 分类映射 ──────────────────────────────────────────────
FACTOR_CATEGORIES: dict[str, list[str]] = {
"动量": ["momentum_5", "momentum_10", "momentum_20", "momentum_60"],
"RSI": ["rsi_7", "rsi_14"],
"MACD": ["macd", "macd_5_35_5"],
"量价": ["vol_ratio_5", "vol_ratio_20", "vol_chg_5"],
"布林": ["boll", "boll_width"],
"ATR": ["atr_14", "atr_ratio_14"],
"均线": ["ma_cross_5_20", "ma_cross_10_60", "ma_dev_20", "ma_dev_60"],
"波动率": ["volatility_20", "volatility_60", "down_vol_20"],
"换手率": ["turnover_5", "turnover_chg_5"],
"振幅": ["amplitude_5", "amplitude_20"],
"基本面": ["roe", "pe", "pb", "ep"],
"情绪": ["news_sent_5", "news_sent_20", "news_conf_5", "sent_delta_5"],
}
def get_factor(name: str, **overrides) -> BaseFactor:
"""按名称获取因子实例。
参数:
name: 因子名称 'momentum_20'
**overrides: 覆盖默认参数
返回:
BaseFactor 实例
"""
if name not in _BUILTIN_FACTORIES:
raise KeyError(f"未知因子: '{name}'。可用: {list(_BUILTIN_FACTORIES)}")
factor = _BUILTIN_FACTORIES[name]()
if overrides:
for k, v in overrides.items():
if hasattr(factor, k):
setattr(factor, k, v)
# 更新 factor.name
if hasattr(factor, "name"):
factor.name = name
return factor
def list_factors(category: str | None = None) -> list[str]:
"""列出所有可用因子名称。"""
if category and category in FACTOR_CATEGORIES:
return FACTOR_CATEGORIES[category]
return list(_BUILTIN_FACTORIES)
def list_categories() -> list[str]:
"""列出所有因子分类。"""
return list(FACTOR_CATEGORIES)
+373
View File
@@ -0,0 +1,373 @@
"""
新闻公告数据源
支持三种数据来源
1. AkShare stock_news_em东方财富个股新闻
2. MariaDB xwlb_daily_ext新闻联播分割数据
3. MCP trendradar-news外部新闻聚合服务
统一返回格式DataFrame (date, title, content, source, url)
"""
import json
import os
import time
from datetime import datetime, timedelta
import akshare as ak
import pandas as pd
import requests
# 确保 .env 已加载
import config.settings # noqa: F401
class NewsSource:
"""
新闻数据源聚合器
参数:
use_akshare: 启用 AkShare 新闻接口
use_xwlb: 启用新闻联播数据库
use_mcp: 启用 MCP 新闻服务
mcp_url: MCP 服务地址
request_delay: 请求间隔避免被限频
"""
def __init__(
self,
use_akshare: bool = True,
use_xwlb: bool = True,
use_mcp: bool = True,
mcp_url: str | None = None,
request_delay: float = 1.0,
):
self.use_akshare = use_akshare
self.use_xwlb = use_xwlb
self.use_mcp = use_mcp
self.mcp_url = mcp_url or os.getenv("NEWS_MCP_URL", "http://192.168.1.160:3333/mcp")
self.request_delay = request_delay
self._mcp_session_id: str | None = None
# ── 统一入口 ──────────────────────────────────────────
def fetch(
self,
ts_code: str,
start: str | None = None,
end: str | None = None,
max_news: int | None = None,
) -> pd.DataFrame:
"""
获取个股相关新闻聚合多源
参数:
ts_code: '000001.SZ'
start: 起始日期 'YYYYMMDD'默认 30 天前
end: 结束日期 'YYYYMMDD'默认今天
max_news: 最多返回条数默认从 .env SENTIMENT_MAX_NEWS_PER_STOCK 读取fallback 30
返回:
DataFrame (date, title, content, source, url)
"""
if max_news is None:
max_news = int(os.getenv("SENTIMENT_MAX_NEWS_PER_STOCK", "30"))
if end is None:
end = time.strftime("%Y%m%d")
if start is None:
start = (datetime.strptime(end, "%Y%m%d") - timedelta(days=30)).strftime("%Y%m%d")
frames = []
if self.use_akshare:
try:
df = self._fetch_akshare(ts_code, start, end)
if not df.empty:
frames.append(df)
time.sleep(self.request_delay)
except Exception as e:
print(" [WARN] AkShare 新闻获取失败 ({}): {}".format(ts_code, e))
if self.use_xwlb:
try:
# xwlb: news_date 是播出日,+1day = 交易日
# 所以查询时 start/end 各减 1 天
xwlb_start = (datetime.strptime(start, "%Y%m%d") - timedelta(days=1)).strftime("%Y%m%d")
xwlb_end = (datetime.strptime(end, "%Y%m%d") - timedelta(days=1)).strftime("%Y%m%d")
df = self._fetch_xwlb(xwlb_start, xwlb_end)
if not df.empty:
frames.append(df)
except Exception as e:
print(" [WARN] xwlb 新闻获取失败: {}".format(e))
if self.use_mcp:
try:
df = self._fetch_mcp(ts_code, start, end)
if not df.empty:
frames.append(df)
except Exception as e:
print(" [WARN] MCP 新闻获取失败: {}".format(e))
if not frames:
return pd.DataFrame(columns=["date", "title", "content", "source", "url"])
result = pd.concat(frames, ignore_index=True)
result = result.drop_duplicates(subset=["title", "date"])
result = result.sort_values("date", ascending=False)
if len(result) > max_news:
result = result.head(max_news)
return result.reset_index(drop=True)
# ── AkShare 数据源 ────────────────────────────────────
def _fetch_akshare(self, ts_code: str, start: str = "", end: str = "") -> pd.DataFrame:
"""
东方财富个股新闻
stock_news_em 返回最新约 10 无日期筛选
如需更多应使用 InfoMoney 或自建爬虫
"""
symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
try:
df = ak.stock_news_em(symbol=symbol.zfill(6))
except Exception:
return pd.DataFrame()
if df is None or df.empty:
return pd.DataFrame()
df = df.rename(columns={
"新闻标题": "title",
"新闻内容": "content",
"发布时间": "date",
"文章来源": "source",
"新闻链接": "url",
})
df["date"] = pd.to_datetime(df["date"]).dt.strftime("%Y%m%d")
df["source"] = "akshare_" + df["source"].fillna("unknown")
# 日期筛选
if start:
df = df[df["date"] >= start]
if end:
df = df[df["date"] <= end]
return df[["date", "title", "content", "source", "url"]]
# ── 新闻联播数据源 ────────────────────────────────────
def _fetch_xwlb(self, start: str, end: str) -> pd.DataFrame:
"""
MariaDB xwlb_daily_ext 表获取新闻联播分割数据
news_date 是播出日期+1day 后成为对市场产生影响的交易日
调用方应传入 (target_start-1, target_end-1)
连接失败时 database/connection.py 自动重建 SSH 隧道
"""
from database.connection import get_engine
start_fmt = "{}-{}-{}".format(start[:4], start[4:6], start[6:8])
end_fmt = "{}-{}-{}".format(end[:4], end[4:6], end[6:8])
day_span = (datetime.strptime(end, "%Y%m%d") - datetime.strptime(start, "%Y%m%d")).days
row_limit = max(100, (day_span + 1) * 40)
sql = """
SELECT news_date, news_title, news_content
FROM xwlb_daily_ext
WHERE news_date BETWEEN %(start)s AND %(end)s
ORDER BY news_date DESC
LIMIT {}
""".format(row_limit)
try:
df = pd.read_sql(
sql, get_engine(),
params={"start": start_fmt, "end": end_fmt},
)
if df.empty:
return pd.DataFrame()
df = df.rename(columns={
"news_date": "date",
"news_title": "title",
"news_content": "content",
})
df["date"] = (
pd.to_datetime(df["date"]) + pd.Timedelta(days=1)
).dt.strftime("%Y%m%d")
df["source"] = "xwlb"
df["url"] = ""
return df[["date", "title", "content", "source", "url"]]
except Exception:
return pd.DataFrame()
# ── MCP 数据源 ────────────────────────────────────────
def _fetch_mcp(self, ts_code: str, start: str, end: str) -> pd.DataFrame:
"""通过 MCP streamable HTTP 调用 trendradar-news 服务。"""
session_id = self._mcp_initialize()
if not session_id:
return pd.DataFrame()
symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
try:
result = self._mcp_call_tool(
session_id,
"search_news",
{"keyword": symbol, "start_date": start, "end_date": end, "limit": 30},
)
if not result:
return pd.DataFrame()
records = self._parse_mcp_result(result)
if not records:
return pd.DataFrame()
df = pd.DataFrame(records)
df["source"] = "mcp_trendradar"
return df[["date", "title", "content", "source", "url"]]
except Exception:
return pd.DataFrame()
def _mcp_initialize(self) -> str | None:
"""MCP streamable HTTP 初始化,获取 session ID。"""
if self._mcp_session_id:
return self._mcp_session_id
try:
resp = requests.post(
self.mcp_url,
json={
"jsonrpc": "2.0",
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": {"name": "cc-cursor", "version": "1.0"},
},
"id": 1,
},
headers={
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
},
timeout=10,
)
# headers 是大小写不敏感的,尝试多种格式
session_id = (
resp.headers.get("mcp-session-id")
or resp.headers.get("Mcp-Session-Id")
or resp.headers.get("MCP-Session-Id")
)
if session_id:
self._mcp_session_id = session_id
else:
print(" [WARN] MCP initialize 未返回 session-id")
return session_id
except Exception as e:
print(" [WARN] MCP 连接失败: {}".format(e))
return None
def _mcp_call_tool(
self, session_id: str, tool_name: str, arguments: dict
) -> dict | None:
"""调用 MCP 工具。"""
try:
resp = requests.post(
self.mcp_url,
json={
"jsonrpc": "2.0",
"method": "tools/call",
"params": {"name": tool_name, "arguments": arguments},
"id": 2,
},
headers={
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
"mcp-session-id": session_id,
},
timeout=15,
)
# 解析 SSE 响应
for line in resp.text.split("\n"):
if line.startswith("data:"):
data = json.loads(line[5:].strip())
return data.get("result", {})
return None
except Exception:
return None
@staticmethod
def _parse_mcp_result(result: dict) -> list[dict]:
"""解析 MCP 返回内容为统一格式。"""
records = []
content = result.get("content", [])
for item in content:
text = item.get("text", "") if isinstance(item, dict) else str(item)
if text:
try:
parsed = json.loads(text)
if isinstance(parsed, list):
for p in parsed:
records.append({
"date": str(p.get("date", p.get("publish_date", ""))),
"title": str(p.get("title", "")),
"content": str(p.get("content", p.get("summary", ""))),
"url": str(p.get("url", "")),
})
elif isinstance(parsed, dict):
records.append({
"date": str(parsed.get("date", parsed.get("publish_date", ""))),
"title": str(parsed.get("title", "")),
"content": str(parsed.get("content", parsed.get("summary", ""))),
"url": str(parsed.get("url", "")),
})
except (json.JSONDecodeError, TypeError):
records.append({
"date": time.strftime("%Y%m%d"),
"title": text[:100],
"content": text,
"url": "",
})
return records
# ── 日期对齐工具 ──────────────────────────────────────────
def align_news_to_trading_days(
news_df: pd.DataFrame,
trading_calendar: pd.DatetimeIndex,
) -> pd.DataFrame:
"""
将新闻日期对齐到最近交易日
非交易日的新闻归入最近的下一个交易日
周六新闻 下周一交易日的因子值中体现
参数:
news_df: 新闻 DataFrame 'date' (YYYYMMDD)
trading_calendar: 交易日 DatetimeIndex
返回:
news_df'date' 列已替换为对齐后的交易日
"""
calendar = trading_calendar.sort_values()
news_dates = pd.to_datetime(news_df["date"], format="%Y%m%d")
aligned = []
for nd in news_dates:
future_dates = calendar[calendar >= nd]
if len(future_dates) > 0:
aligned.append(future_dates[0])
else:
aligned.append(calendar[-1])
news_df = news_df.copy()
news_df["date"] = pd.DatetimeIndex(aligned).strftime("%Y%m%d")
return news_df
+229
View File
@@ -0,0 +1,229 @@
"""
Qwen API 封装DashScope / 本地 Ollama
文本 情绪分析 标准化 JSON 输出
"""
import json
import os
import time
from typing import Optional
import requests
# 确保 .env 已加载(无论从哪个路径导入)
import config.settings # noqa: F401
class QwenClient:
"""
Qwen 情绪分析客户端
支持两种后端
1. DashScope API- 需要 QWEN_API_KEY
2. 本地 Ollama本地- 需要 QWEN_LOCAL_BASE_URL
优先级本地 > API
参数:
api_key: DashScope API Key默认从环境变量 QWEN_API_KEY 读取
model: 模型名称默认从 QWEN_MODEL 读取
local_base_url: Ollama 地址默认从 QWEN_LOCAL_BASE_URL 读取
local_model: Ollama 模型名默认从 QWEN_LOCAL_MODEL 读取
max_retries: 失败重试次数
cache_enabled: 是否启用结果缓存按文本 hash
"""
# 情绪分析 Prompt
SYSTEM_PROMPT = """你是一个专业的金融情绪分析专家。分析以下 A 股相关的新闻/公告文本,返回 JSON 格式的分析结果。
分析维度
1. sentiment_score: -1.0(极度利空) ~ 0(中性) ~ +1.0(极度利好)保留1位小数
2. impact_duration: "short"(1-3个交易日) | "medium"(1-2) | "long"(1个月以上)
3. confidence: 0.0~1.0 置信度
4. key_topics: 涉及的关键主题列表最多5个
5. affected_factors: 可能影响的基本面/技术面因子类型列表
重要规则
- 只考虑对股价的直接影响
- 中性新闻常规报道例行公告 0
- 重大利好业绩超预期政策支持大订单 >0.5
- 重大利空亏损处罚减持诉讼 <-0.5
- 如果不确定confidences 要低 (<0.5)
返回格式 JSON
{"sentiment_score": 0.0, "impact_duration": "short", "confidence": 0.5, "key_topics": [], "affected_factors": []}"""
def __init__(
self,
api_key: str | None = None,
model: str | None = None,
local_base_url: str | None = None,
local_model: str | None = None,
max_retries: int = 3,
cache_enabled: bool = True,
):
self.api_key = api_key or os.getenv("QWEN_API_KEY", "")
self.model = model or os.getenv("QWEN_MODEL", "qwen-turbo")
self.local_base_url = local_base_url or os.getenv("QWEN_LOCAL_BASE_URL", "")
self.local_model = local_model or os.getenv("QWEN_LOCAL_MODEL", "qwen2.5:7b")
self.max_retries = max_retries
self.cache_enabled = cache_enabled
self._cache: dict[str, dict] = {}
self._is_local = bool(self.local_base_url)
# ── 单条分析 ──────────────────────────────────────────
def analyze_sentiment(self, text: str) -> dict:
"""
分析单条文本的情绪
返回:
{"sentiment_score": float, "impact_duration": str, "confidence": float,
"key_topics": list, "affected_factors": list}
"""
if not text or not text.strip():
return self._empty_result()
cache_key = str(hash(text))
if self.cache_enabled and cache_key in self._cache:
return self._cache[cache_key]
for attempt in range(self.max_retries):
try:
if self._is_local:
result = self._call_ollama(text)
else:
result = self._call_dashscope(text)
if result and "sentiment_score" in result:
if self.cache_enabled:
self._cache[cache_key] = result
return result
except Exception as e:
if attempt < self.max_retries - 1:
time.sleep(1 * (attempt + 1))
else:
print(f" [ERROR] Qwen 分析失败(重试{self.max_retries}次): {e}")
return self._empty_result()
# ── 批量分析 ──────────────────────────────────────────
def analyze_batch(
self,
texts: list[str],
batch_size: int = 10,
progress_callback: callable = None, # type: ignore
) -> list[dict]:
"""
批量分析逐条调用节省并发成本
参数:
texts: 文本列表
batch_size: 每批次数量
progress_callback: 进度回调 fn(done, total)
返回:
[{"sentiment_score": ..., ...}, ...]
"""
results = []
total = len(texts)
for i, text in enumerate(texts):
result = self.analyze_sentiment(text)
results.append(result)
if progress_callback:
progress_callback(i + 1, total)
# 批次间短暂休息,避免触发限频
if (i + 1) % batch_size == 0 and i < total - 1:
time.sleep(0.5)
return results
# ── DashScope API ──────────────────────────────────────
def _call_dashscope(self, text: str) -> dict:
"""调用 DashScope API。"""
resp = requests.post(
"https://dashscope.aliyuncs.com/compatible-mode/v1/chat/completions",
headers={
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
},
json={
"model": self.model,
"messages": [
{"role": "system", "content": self.SYSTEM_PROMPT},
{"role": "user", "content": text[:4000]}, # 截断长文本
],
"temperature": 0.1,
"max_tokens": 500,
},
timeout=30,
)
resp.raise_for_status()
body = resp.json()
content = body["choices"][0]["message"]["content"]
return self._parse_json(content)
# ── Ollama 本地调用 ────────────────────────────────────
def _call_ollama(self, text: str) -> dict:
"""调用本地 Ollama。"""
resp = requests.post(
f"{self.local_base_url.rstrip('/')}/chat/completions",
json={
"model": self.local_model,
"messages": [
{"role": "system", "content": self.SYSTEM_PROMPT},
{"role": "user", "content": text[:4000]},
],
"temperature": 0.1,
"max_tokens": 500,
},
timeout=60,
)
resp.raise_for_status()
body = resp.json()
content = body["choices"][0]["message"]["content"]
return self._parse_json(content)
# ── JSON 解析 ──────────────────────────────────────────
@staticmethod
def _parse_json(content: str) -> dict:
"""从模型输出中提取 JSON。"""
# 尝试直接解析
try:
return json.loads(content)
except json.JSONDecodeError:
pass
# 尝试提取 ```json ... ``` 代码块
if "```json" in content:
start = content.find("```json") + 7
end = content.find("```", start)
if end > start:
try:
return json.loads(content[start:end].strip())
except json.JSONDecodeError:
pass
# 尝试找 { ... }
brace_start = content.find("{")
brace_end = content.rfind("}")
if brace_start >= 0 and brace_end > brace_start:
try:
return json.loads(content[brace_start:brace_end + 1])
except json.JSONDecodeError:
pass
return QwenClient._empty_result()
@staticmethod
def _empty_result() -> dict:
return {
"sentiment_score": 0.0,
"impact_duration": "short",
"confidence": 0.0,
"key_topics": [],
"affected_factors": [],
}
@@ -0,0 +1,311 @@
"""
情绪因子计算引擎
全链路数据获取 Qwen 分析 因子计算 缓存管理
"""
import json
import os
import time
from datetime import datetime
from pathlib import Path
import pandas as pd
from data.data_manager import DataManager
from factors.sentiment.news_source import NewsSource, align_news_to_trading_days
from factors.sentiment.qwen_client import QwenClient
from factors.sentiment.sentiment_factor import (
NewsSentimentFactor,
SentimentMomentumFactor,
SentimentConfidenceFactor,
)
def _get_or_fetch_price(dm, ts_code, fallback_code="000001.SZ"):
"""
获取交易日历价格 DataFrame
1. 优先从 DB 缓存读取目标股票
2. 无缓存则 sync_daily 补齐
3. 补齐失败则 fallback 到备用股票仅用作交易日历不影响新闻分析
"""
import pandas as pd
from database.dao import get_latest_trade_date
# 检查 DB 缓存
if get_latest_trade_date(ts_code):
daily = dm.get_daily(ts_code)
if daily is not None and not daily.empty:
return daily.set_index("trade_date").sort_index()
# 无缓存 → 尝试补齐
print(" [SentimentEngine] {} 无 DB 缓存,尝试 sync_daily...".format(ts_code))
try:
n = dm.sync_daily(ts_code)
if n > 0:
daily = dm.get_daily(ts_code)
if daily is not None and not daily.empty:
return daily.set_index("trade_date").sort_index()
except Exception as e:
print(" [SentimentEngine] sync_daily 失败: {}".format(e))
# Fallback
print(" [SentimentEngine] 使用 {} 交易日历作为 fallback".format(fallback_code))
try:
fallback = dm.get_daily(fallback_code)
if fallback is not None and not fallback.empty:
return fallback.set_index("trade_date").sort_index()
except Exception:
pass
return None
class SentimentEngine:
"""
情绪因子计算引擎
参数:
data_manager: DataManager 实例
qwen_client: QwenClient 实例可选默认从环境变量创建
news_source: NewsSource 实例可选
cache_dir: 缓存目录
"""
def __init__(
self,
data_manager: DataManager,
qwen_client: QwenClient | None = None,
news_source: NewsSource | None = None,
cache_dir: str | None = None,
):
self._dm = data_manager
self._qwen = qwen_client or QwenClient()
self._news_source = news_source or NewsSource()
self._cache_dir = Path(cache_dir or os.path.join(
os.path.dirname(__file__), "..", "..", ".cache", "sentiment"
))
self._cache_dir.mkdir(parents=True, exist_ok=True)
# ── 核心方法 ──────────────────────────────────────────
def compute(
self,
ts_code: str,
start: str | None = None,
end: str | None = None,
max_news: int = 20,
) -> pd.DataFrame:
"""
计算单只股票的情绪因子含缓存
返回:
DataFrame (trade_date, news_sent_5, news_conf_5, sent_delta_5)
"""
from database.dao import get_latest_trade_date
# 1. 获取交易日历:优先目标股票 DB 缓存,无则尝试补齐,失败则 fallback
price = _get_or_fetch_price(self._dm, ts_code)
if price is None:
return pd.DataFrame()
daily_idx = pd.to_datetime(price.index, format="%Y%m%d", errors="coerce")
# 2. 检查缓存
cache_key = "sent_{}_{}_{}".format(ts_code, start or "all", end or "now")
cached = self._load_cache(cache_key)
if cached is not None:
return cached
# 3. 获取新闻
max_n = max_news or int(os.getenv("SENTIMENT_MAX_NEWS_PER_STOCK", "30"))
news_df = self._news_source.fetch(ts_code, start=start, end=end, max_news=max_n)
# 4. 对齐到交易日
if not news_df.empty:
news_df = align_news_to_trading_days(news_df, daily_idx)
# 5. Qwen 情绪分析(有 API key 时才执行)
sentiment_df = self._analyze_news(news_df)
# 6. 计算因子
factor_dfs = {}
if not sentiment_df.empty:
for factor_cls, kwargs in [
(NewsSentimentFactor, {"window": 5, "sentiment_df": sentiment_df}),
(SentimentConfidenceFactor, {"window": 5, "sentiment_df": sentiment_df}),
(SentimentMomentumFactor, {"period": 5, "sentiment_df": sentiment_df}),
]:
f = factor_cls(**kwargs)
series = f.calculate(price)
factor_dfs[f.name] = series
if factor_dfs:
result = pd.DataFrame(factor_dfs)
else:
result = pd.DataFrame(index=price.index)
# 7. 缓存
self._save_cache(cache_key, result)
return result
def compute_batch(
self,
ts_codes: list[str],
date: str | None = None,
max_news: int = 20,
) -> dict[str, pd.DataFrame]:
"""
批量计算多只股票的情绪因子
返回:
{ts_code: factor_df}
"""
results = {}
total = len(ts_codes)
for i, ts_code in enumerate(ts_codes):
try:
results[ts_code] = self.compute(
ts_code, start=None if date is None else date, end=date,
max_news=max_news,
)
except Exception as e:
print(f"[WARN] {ts_code} 情绪因子失败: {e}")
if (i + 1) % 10 == 0:
print(f"[SentimentEngine] {i + 1}/{total}")
return results
# ── 分析范围解析 ──────────────────────────────────────
def get_scope_stocks(self) -> list[str]:
"""
根据 .env 配置解析分析范围
优先级: custom > index > sector > all
"""
scope_type = os.getenv("SENTIMENT_SCOPE_TYPE", "index")
all_stocks = self._dm.get_stock_list()
if scope_type == "custom":
codes = os.getenv("SENTIMENT_SCOPE_CUSTOM", "")
return [c.strip() for c in codes.split(",") if c.strip()]
if scope_type == "sector":
return self._get_sector_stocks(all_stocks)
if scope_type == "index":
return self._get_index_stocks()
# all
return list(all_stocks.index)
def _get_index_stocks(self) -> list[str]:
"""根据指数成分股获取股票列表。"""
indexes_str = os.getenv("SENTIMENT_SCOPE_INDEXES", "000300")
indexes = [i.strip() for i in indexes_str.split(",") if i.strip()]
stocks = set()
import akshare as ak
for idx_code in indexes:
try:
df = ak.index_stock_cons(symbol=idx_code)
if df is not None and not df.empty:
# 列名可能是 '品种代码' 或 'stock_code'
col = next((c for c in df.columns if "代码" in c or "code" in c.lower()), df.columns[0])
for code in df[col]:
stocks.add(f"{code}.SZ" if code.startswith("0") or code.startswith("3") else f"{code}.SH")
except Exception as e:
print(f" [WARN] 指数 {idx_code} 成分股获取失败: {e}")
return list(stocks)
def _get_sector_stocks(self, all_stocks: pd.DataFrame) -> list[str]:
"""根据板块名称筛选股票。"""
sectors_str = os.getenv("SENTIMENT_SCOPE_SECTORS", "")
sectors = [s.strip() for s in sectors_str.split(",") if s.strip()]
if not sectors:
return list(all_stocks.index)
# 利用 AkShare 行业分类筛选
import akshare as ak
all_industry = ak.stock_board_industry_name_em()
matched = all_industry[all_industry["板块名称"].isin(sectors)]
stocks = set()
for _, row in matched.iterrows():
try:
df = ak.stock_board_industry_cons_em(symbol=row["板块名称"])
if df is not None and not df.empty:
code_col = next((c for c in df.columns if "代码" in c), df.columns[0])
for code in df[code_col]:
stocks.add(code)
except Exception:
continue
return list(stocks)
# ── 新闻情绪分析 ──────────────────────────────────────
def _analyze_news(self, news_df: pd.DataFrame) -> pd.DataFrame:
"""
对新闻列表执行情绪分析
返回:
DataFrame (date, sentiment_score, confidence, impact_duration, key_topics, title, content)
"""
if news_df.empty:
return pd.DataFrame()
# 检查是否有 Qwen API 配置
if not self._qwen.api_key and not self._qwen.local_base_url:
return pd.DataFrame()
results = []
for _, row in news_df.iterrows():
title = row.get("title", "")
content = row.get("content", "")
# 优先用标题+内容,文本过短时只用标题
text = f"{title}\n{content}" if len(str(content)) > 20 else title
if len(text.strip()) < 10:
continue
analysis = self._qwen.analyze_sentiment(text)
results.append({
"date": row["date"],
"title": title,
"sentiment_score": analysis.get("sentiment_score", 0.0),
"confidence": analysis.get("confidence", 0.0),
"impact_duration": analysis.get("impact_duration", "short"),
"key_topics": json.dumps(analysis.get("key_topics", [])),
})
if not results:
return pd.DataFrame()
return pd.DataFrame(results)
# ── 缓存管理 ──────────────────────────────────────────
def _load_cache(self, key: str) -> pd.DataFrame | None:
"""加载缓存。缓存有效期 6 小时。"""
cache_file = self._cache_dir / f"{key.replace('/', '_')}.parquet"
if not cache_file.exists():
return None
mtime = cache_file.stat().st_mtime
if time.time() - mtime > 6 * 3600:
return None
try:
return pd.read_parquet(cache_file)
except Exception:
return None
def _save_cache(self, key: str, df: pd.DataFrame) -> None:
cache_file = self._cache_dir / f"{key.replace('/', '_')}.parquet"
try:
df.to_parquet(cache_file)
except Exception:
pass
def clear_cache(self) -> int:
"""清空缓存,返回删除文件数。"""
count = 0
for f in self._cache_dir.glob("*.parquet"):
f.unlink()
count += 1
return count
@@ -0,0 +1,168 @@
"""
情绪因子
Qwen 输出的情绪分数转换为量化因子值
"""
import numpy as np
import pandas as pd
from factors.base import BaseFactor
class NewsSentimentFactor(BaseFactor):
"""
新闻情绪因子
将多条新闻的情绪分数按时间加权聚合到每个交易日
参数:
window: 滚动窗口交易日
decay: 指数衰减系数越大衰减越快0 表示等权
sentiment_df: 情绪分析结果 DataFrame
(date, sentiment_score, confidence, title, content)
"""
category = "sentiment"
def __init__(self, window: int = 5, decay: float = 0.3, sentiment_df: pd.DataFrame | None = None):
self.window = window
self.decay = decay
self.sentiment_df = sentiment_df
self.name = f"news_sent_{window}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
if self.sentiment_df is None or self.sentiment_df.empty:
return pd.Series(float("nan"), index=df.index, name=self.name)
return _aggregate_sentiment(
df, self.sentiment_df, self.window, self.decay
)
def get_required_columns(self) -> list[str]:
return []
class SentimentMomentumFactor(BaseFactor):
"""
情绪动量因子
当前情绪 - N 日前情绪衡量情绪变化方向
"""
category = "sentiment"
def __init__(self, period: int = 5, sentiment_df: pd.DataFrame | None = None):
self.period = period
self.sentiment_df = sentiment_df
self.name = f"sent_delta_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
base = NewsSentimentFactor(window=1, sentiment_df=self.sentiment_df).calculate(df)
return base.diff(self.period)
def get_required_columns(self) -> list[str]:
return []
class SentimentConfidenceFactor(BaseFactor):
"""
情绪置信度因子
新闻情绪分析的置信度越高因子绝对值越大方向同 sentiment_score
sentiment_score × confidence 高置信利好=正值大高置信利空=负值大
"""
category = "sentiment"
def __init__(self, window: int = 5, sentiment_df: pd.DataFrame | None = None):
self.window = window
self.sentiment_df = sentiment_df
self.name = f"news_conf_{window}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
if self.sentiment_df is None or self.sentiment_df.empty:
return pd.Series(float("nan"), index=df.index, name=self.name)
sdf = self.sentiment_df.copy()
# 设置加权分数
if "confidence" in sdf.columns and "sentiment_score" in sdf.columns:
sdf["weighted_score"] = sdf["sentiment_score"] * sdf["confidence"]
else:
return pd.Series(float("nan"), index=df.index, name=self.name)
return _aggregate_sentiment(df, sdf, self.window, decay=0.3, score_col="weighted_score")
def get_required_columns(self) -> list[str]:
return []
# ── 情绪聚合工具函数 ──────────────────────────────────────
def _aggregate_sentiment(
daily_df: pd.DataFrame,
sentiment_df: pd.DataFrame,
window: int,
decay: float,
score_col: str = "sentiment_score",
) -> pd.Series:
"""
将情绪分数按时间加权聚合到交易日
逻辑
1. 对每个交易日 t找到 [t - window + 1, t] 范围内的所有新闻
2. time_decay = exp(-decay * days_from_t) 加权
3. confidence如有加权
4. 返回加权平均情绪分数
参数:
daily_df: 日线 DataFrame提供 index 和日期对齐
sentiment_df: 情绪 DataFramedate + score_col
window: 窗口大小
decay: 衰减系数
score_col: 情绪分数列名
"""
if sentiment_df.empty:
return pd.Series(float("nan"), index=daily_df.index, name=score_col)
# 统一日期格式
sdf = sentiment_df.copy()
sdf["date"] = pd.to_datetime(sdf["date"], format="%Y%m%d", errors="coerce")
sdf = sdf.dropna(subset=["date"])
sdf = sdf.sort_values("date")
daily_idx = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce")
if daily_idx.isna().all():
daily_idx = pd.to_datetime(daily_df.index)
result = pd.Series(float("nan"), index=daily_df.index)
# 对 news 日期建立搜索索引
news_dates = sdf["date"].values
for i, dt in enumerate(daily_idx):
if pd.isna(dt):
continue
# 窗口起始
window_start = dt - pd.Timedelta(days=window * 2) # 宽窗覆盖非交易日
mask = (news_dates >= window_start) & (news_dates <= dt)
candidates = sdf[mask]
if candidates.empty:
continue
# 时间衰减权重
days_diff = (dt - candidates["date"]).dt.days
time_weights = np.exp(-decay * days_diff)
# 置信度权重(如有)
conf_weights = candidates.get("confidence", pd.Series(1.0, index=candidates.index)).fillna(0.5)
scores = candidates[score_col].fillna(0.0)
total_weight = (time_weights * conf_weights).sum()
if total_weight > 0:
result.iloc[i] = (scores * time_weights * conf_weights).sum() / total_weight
result.name = score_col
return result.astype("float64")
+24
View File
@@ -0,0 +1,24 @@
"""
振幅因子
"""
import pandas as pd
from factors.base import BaseFactor
class AmplitudeFactor(BaseFactor):
"""N 日均振幅 = mean((high - low) / close, N) * 100"""
category = "technical"
def __init__(self, period: int = 5):
self.period = period
self.name = f"amplitude_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
daily_amp = (df["high"] - df["low"]) / df["close"].replace(0, float("nan")) * 100
return daily_amp.rolling(window=self.period, min_periods=self.period).mean()
def get_required_columns(self) -> list[str]:
return ["high", "low", "close"]
+47
View File
@@ -0,0 +1,47 @@
"""
ATR 平均真实波幅因子
"""
import pandas as pd
from factors.base import BaseFactor
class ATRFactor(BaseFactor):
"""Average True Range,衡量波动性。"""
category = "technical"
def __init__(self, period: int = 14):
self.period = period
self.name = f"atr_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
high, low, close = df["high"], df["low"], df["close"]
prev_close = close.shift(1)
tr = pd.concat([
(high - low).abs(),
(high - prev_close).abs(),
(low - prev_close).abs(),
], axis=1).max(axis=1)
return tr.ewm(span=self.period, min_periods=self.period).mean()
def get_required_columns(self) -> list[str]:
return ["high", "low", "close"]
class ATRRatioFactor(BaseFactor):
"""ATR / close 归一化,便于跨股票比较。"""
category = "technical"
def __init__(self, period: int = 14):
self.period = period
self.name = f"atr_ratio_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
atr = ATRFactor(period=self.period).calculate(df)
return atr / df["close"].replace(0, float("nan")) * 100
def get_required_columns(self) -> list[str]:
return ["high", "low", "close"]
+53
View File
@@ -0,0 +1,53 @@
"""
布林带因子价格在布林带中的位置
"""
import pandas as pd
from factors.base import BaseFactor
class BollingerFactor(BaseFactor):
"""
布林带位置 = (close - middle) / (upper - lower)
值在 0~1 之间接近 0 表示在下轨接近 1 表示在上轨
"""
category = "technical"
def __init__(self, period: int = 20, std: float = 2.0):
self.period = period
self.std = std
self.name = f"boll_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
middle = df["close"].rolling(window=self.period, min_periods=self.period).mean()
std = df["close"].rolling(window=self.period, min_periods=self.period).std()
upper = middle + self.std * std
lower = middle - self.std * std
band_width = upper - lower
return ((df["close"] - lower) / band_width.replace(0, float("nan"))).clip(0, 1)
def get_required_columns(self) -> list[str]:
return ["close"]
class BollingerWidthFactor(BaseFactor):
"""布林带宽度 = (upper - lower) / middle * 100"""
category = "technical"
def __init__(self, period: int = 20, std: float = 2.0):
self.period = period
self.std = std
self.name = f"boll_width_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
middle = df["close"].rolling(window=self.period, min_periods=self.period).mean()
std = df["close"].rolling(window=self.period, min_periods=self.period).std()
band_width = 2 * self.std * std
return band_width / middle.replace(0, float("nan")) * 100
def get_required_columns(self) -> list[str]:
return ["close"]
+49
View File
@@ -0,0 +1,49 @@
"""
均线交叉因子
"""
import pandas as pd
from factors.base import BaseFactor
class MACrossFactor(BaseFactor):
"""
均线交叉信号
返回fast_ma / slow_ma - 1正值表示短期均线在上方
"""
category = "technical"
def __init__(self, fast: int = 5, slow: int = 20):
self.fast = fast
self.slow = slow
self.name = f"ma_cross_{fast}_{slow}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
ma_fast = df["close"].rolling(window=self.fast, min_periods=self.fast).mean()
ma_slow = df["close"].rolling(window=self.slow, min_periods=self.slow).mean()
return (ma_fast / ma_slow.replace(0, float("nan"))) - 1
def get_required_columns(self) -> list[str]:
return ["close"]
class MADeviationFactor(BaseFactor):
"""
价格偏离均线程度 = (close - ma) / ma * 100
"""
category = "technical"
def __init__(self, period: int = 20):
self.period = period
self.name = f"ma_dev_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
ma = df["close"].rolling(window=self.period, min_periods=self.period).mean()
return (df["close"] - ma) / ma.replace(0, float("nan")) * 100
def get_required_columns(self) -> list[str]:
return ["close"]
+47
View File
@@ -0,0 +1,47 @@
"""
MACD 因子
"""
import pandas as pd
from factors.base import BaseFactor
class MACDFactor(BaseFactor):
"""
MACD 系列因子
返回 DIF/DEA/HIST 三个值使用 calculate() 返回 HIST
单独方法获取 DIF/DEA
"""
category = "technical"
def __init__(self, fast: int = 12, slow: int = 26, signal: int = 9):
self.fast = fast
self.slow = slow
self.signal = signal
self.name = f"macd_{fast}_{slow}_{signal}"
def _ema(self, series: pd.Series, span: int) -> pd.Series:
return series.ewm(span=span, min_periods=span).mean()
def calculate(self, df: pd.DataFrame) -> pd.Series:
"""返回 MACD 柱(DIF - DEA)。"""
ema_fast = self._ema(df["close"], self.fast)
ema_slow = self._ema(df["close"], self.slow)
dif = ema_fast - ema_slow
dea = self._ema(dif, self.signal)
return dif - dea
def dif(self, df: pd.DataFrame) -> pd.Series:
ema_fast = self._ema(df["close"], self.fast)
ema_slow = self._ema(df["close"], self.slow)
return ema_fast - ema_slow
def dea(self, df: pd.DataFrame) -> pd.Series:
dif = self.dif(df)
return self._ema(dif, self.signal)
def get_required_columns(self) -> list[str]:
return ["close"]
+23
View File
@@ -0,0 +1,23 @@
"""
动量因子N 日收益率
"""
import pandas as pd
from factors.base import BaseFactor
class MomentumFactor(BaseFactor):
"""N 日价格动量 = (close_t - close_{t-N}) / close_{t-N} * 100"""
category = "technical"
def __init__(self, period: int = 20):
self.period = period
self.name = f"momentum_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
return df["close"].pct_change(periods=self.period) * 100
def get_required_columns(self) -> list[str]:
return ["close"]
+29
View File
@@ -0,0 +1,29 @@
"""
RSI 相对强弱因子
"""
import pandas as pd
from factors.base import BaseFactor
class RSIFactor(BaseFactor):
"""Wilder's RSI = 100 - 100 / (1 + RS), RS = avg_gain / avg_loss"""
category = "technical"
def __init__(self, period: int = 14):
self.period = period
self.name = f"rsi_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
delta = df["close"].diff()
gain = delta.clip(lower=0)
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)
def get_required_columns(self) -> list[str]:
return ["close"]
+40
View File
@@ -0,0 +1,40 @@
"""
换手率因子
"""
import pandas as pd
from factors.base import BaseFactor
class TurnoverFactor(BaseFactor):
"""N 日均换手率。"""
category = "technical"
def __init__(self, period: int = 5):
self.period = period
self.name = f"turnover_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
return df["turnover_rate"].rolling(window=self.period, min_periods=self.period).mean()
def get_required_columns(self) -> list[str]:
return ["turnover_rate"]
class TurnoverChangeFactor(BaseFactor):
"""换手率变化 = 当日换手率 / N 日均换手率。"""
category = "technical"
def __init__(self, period: int = 5):
self.period = period
self.name = f"turnover_chg_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
avg = df["turnover_rate"].rolling(window=self.period, min_periods=self.period).mean()
return df["turnover_rate"] / avg.replace(0, float("nan"))
def get_required_columns(self) -> list[str]:
return ["turnover_rate"]
+42
View File
@@ -0,0 +1,42 @@
"""
波动率因子
"""
import pandas as pd
from factors.base import BaseFactor
class VolatilityFactor(BaseFactor):
"""N 日年化波动率 = std(daily_return, N) * sqrt(252) * 100"""
category = "technical"
def __init__(self, period: int = 20):
self.period = period
self.name = f"volatility_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
daily_ret = df["close"].pct_change()
return daily_ret.rolling(window=self.period, min_periods=self.period).std() * (252 ** 0.5) * 100
def get_required_columns(self) -> list[str]:
return ["close"]
class DownsideVolatilityFactor(BaseFactor):
"""下行波动率:只计算负收益的标准差。"""
category = "technical"
def __init__(self, period: int = 20):
self.period = period
self.name = f"down_vol_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
daily_ret = df["close"].pct_change()
downside = daily_ret.clip(upper=0)
return downside.rolling(window=self.period, min_periods=self.period).std() * (252 ** 0.5) * 100
def get_required_columns(self) -> list[str]:
return ["close"]
+40
View File
@@ -0,0 +1,40 @@
"""
量价因子量比 / 成交量变化率
"""
import pandas as pd
from factors.base import BaseFactor
class VolumeFactor(BaseFactor):
"""N 日均量比 = vol / mean(vol, N)"""
category = "technical"
def __init__(self, period: int = 5):
self.period = period
self.name = f"vol_ratio_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
avg_vol = df["vol"].rolling(window=self.period, min_periods=self.period).mean()
return df["vol"] / avg_vol.replace(0, float("nan"))
def get_required_columns(self) -> list[str]:
return ["vol"]
class VolumeChangeFactor(BaseFactor):
"""成交量 N 日变化率"""
category = "technical"
def __init__(self, period: int = 5):
self.period = period
self.name = f"vol_chg_{period}"
def calculate(self, df: pd.DataFrame) -> pd.Series:
return df["vol"].pct_change(periods=self.period) * 100
def get_required_columns(self) -> list[str]:
return ["vol"]
View File
+124
View File
@@ -0,0 +1,124 @@
"""
ML 模型回测集成
MLStrategy: ML 预测值作为交易信号接入回测引擎
MLBenchmark: 多模型基准对比
"""
import numpy as np
import pandas as pd
from backtest.base import BaseStrategy
from backtest.vectorbt.engine import VectorBTEngine
from models.base import BaseModel
from models.features import FeatureEngine
class MLStrategy(BaseStrategy):
"""
ML 预测 交易信号
用模型预测未来 N 日收益按预测值分位数生成信号
- 预测值 > buy_quantile 买入
- 预测值 < sell_quantile 平仓
参数:
model: 已训练的 BaseModel
feature_engine: fit FeatureEngine
buy_quantile: 买入分位阈值0.7 = 预测值最高的30%买入
sell_quantile: 卖出分位阈值0.3 = 预测值最低的30%平仓
rebalance_freq: 调仓间隔交易日
"""
category = "ml"
def __init__(
self,
model: BaseModel,
feature_engine: FeatureEngine,
buy_quantile: float = 0.7,
sell_quantile: float = 0.3,
rebalance_freq: int = 5,
):
self.model = model
self.feature_engine = feature_engine
self.buy_quantile = buy_quantile
self.sell_quantile = sell_quantile
self.rebalance_freq = rebalance_freq
self.name = f"ml_{model.name}"
def generate_signals(self, factor_df: pd.DataFrame) -> pd.Series:
X, _ = self.feature_engine.build(factor_df, factor_df, fit=False)
if X.empty:
return pd.Series(-1, index=factor_df.index)
preds = self.model.predict(X)
# 用预测值本身的分布作为阈值(相对排序,避免模型偏差影响)
buy_threshold = preds.quantile(self.buy_quantile)
sell_threshold = preds.quantile(self.sell_quantile)
signals = pd.Series(-1, index=factor_df.index)
common = signals.index.intersection(preds.index)
buy_mask = preds.loc[common] > buy_threshold
sell_mask = preds.loc[common] < sell_threshold
signals.loc[buy_mask[buy_mask].index] = 1
signals.loc[sell_mask[sell_mask].index] = 0
signals = self._filter_rebalance(signals)
return signals
def _filter_rebalance(self, signals: pd.Series) -> pd.Series:
"""每隔 rebalance_freq 个交易日保留第一个非持有信号。"""
result = signals.copy()
last_active = -self.rebalance_freq - 1
for i in range(len(result)):
sig = result.iloc[i]
if sig in (0, 1):
if i - last_active >= self.rebalance_freq:
last_active = i
else:
result.iloc[i] = -1
return result
class MLBenchmark:
"""ML 模型基准对比测试。"""
def __init__(
self,
models: list[BaseModel],
feature_engine: FeatureEngine,
price_df: pd.DataFrame,
factor_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.bt_engine = bt_engine or VectorBTEngine()
def run(self) -> pd.DataFrame:
"""对比各模型的预测质量和回测表现。"""
rows = []
for model in self.models:
strategy = MLStrategy(model, self.feature_engine)
report = self.bt_engine.run(strategy, self.price_df, self.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
rows.append({
"model": model.name,
"ic": round(ic, 4),
"total_return": report.total_return,
"cagr": report.cagr,
"max_dd": report.max_drawdown,
"sharpe": report.sharpe_ratio,
"win_rate": report.win_rate,
"trades": report.total_trades,
})
return pd.DataFrame(rows).set_index("model")
+42
View File
@@ -0,0 +1,42 @@
"""
ML 模型抽象基类
统一接口fit(X, y) predict(X) save/load
"""
from abc import ABC, abstractmethod
import pickle
import pandas as pd
class BaseModel(ABC):
"""ML 模型抽象基类。"""
name: str = ""
@abstractmethod
def fit(self, X: pd.DataFrame, y: pd.Series) -> "BaseModel":
"""训练模型。返回 self 支持链式调用。"""
...
@abstractmethod
def predict(self, X: pd.DataFrame) -> pd.Series:
"""返回预测值(回归值)。"""
...
@abstractmethod
def get_feature_importance(self) -> pd.DataFrame:
"""特征重要性 DataFramecolumns=[feature, importance]。"""
...
def save(self, path: str) -> None:
"""保存模型到文件(pickle)。"""
with open(path, "wb") as f:
pickle.dump(self, f)
@classmethod
def load(cls, path: str) -> "BaseModel":
"""从文件加载模型。"""
with open(path, "rb") as f:
return pickle.load(f)
View File
+113
View File
@@ -0,0 +1,113 @@
"""
CatBoost 模型封装
"""
import numpy as np
import pandas as pd
from catboost import CatBoostRegressor, Pool
from models.base import BaseModel
from sklearn.model_selection import TimeSeriesSplit
_DEFAULT_PARAMS = {
"loss_function": "RMSE",
"iterations": 1000,
"learning_rate": 0.03,
"depth": 5,
"random_seed": 42,
"verbose": False,
"allow_writing_files": False,
"min_data_in_leaf": 20,
}
class CatBoostModel(BaseModel):
"""CatBoost 回归模型。"""
name = "catboost"
def __init__(
self,
params: dict | None = None,
early_stopping: int = 50,
eval_ratio: float = 0.2,
random_seed: int = 42,
):
self.params = params or _DEFAULT_PARAMS.copy()
self.params["random_seed"] = random_seed
self.early_stopping = early_stopping
self.eval_ratio = eval_ratio
self._model: CatBoostRegressor | None = None
self._feature_names: list[str] = []
def fit(self, X: pd.DataFrame, y: pd.Series) -> "CatBoostModel":
self._feature_names = list(X.columns)
n = len(X)
val_size = int(n * self.eval_ratio)
X_train, y_train = X, y
eval_set = None
if val_size >= 50:
split_idx = n - val_size
X_train, X_val = X.iloc[:split_idx], X.iloc[split_idx:]
y_train, y_val = y.iloc[:split_idx], y.iloc[split_idx:]
eval_set = Pool(X_val, y_val)
self._model = CatBoostRegressor(**self.params)
self._model.fit(
X_train, y_train,
eval_set=eval_set,
early_stopping_rounds=self.early_stopping if eval_set else None,
verbose=False,
)
return self
def predict(self, X: pd.DataFrame) -> pd.Series:
if self._model is None:
raise RuntimeError("模型尚未训练")
preds = self._model.predict(X[self._feature_names])
return pd.Series(preds, index=X.index, name="pred")
def get_feature_importance(self) -> pd.DataFrame:
if self._model is None:
return pd.DataFrame()
imp = self._model.get_feature_importance()
names = self._feature_names
df = pd.DataFrame({"feature": names, "importance": imp})
total = df["importance"].sum()
df["importance_pct"] = df["importance"] / total * 100 if total > 0 else 0
return df.sort_values("importance", ascending=False)
def cv_evaluate(
self, X: pd.DataFrame, y: pd.Series, n_folds: int = 5
) -> pd.DataFrame:
"""时间序列交叉验证评估。"""
tscv = TimeSeriesSplit(n_splits=n_folds)
results = []
for fold, (train_idx, test_idx) in enumerate(tscv.split(X)):
X_train, X_test = X.iloc[train_idx], X.iloc[test_idx]
y_train, y_test = y.iloc[train_idx], y.iloc[test_idx]
model = CatBoostModel(
params=self.params,
early_stopping=self.early_stopping,
eval_ratio=0.0,
)
model.fit(X_train, y_train)
preds = model.predict(X_test)
ic = preds.corr(y_test)
mse = ((preds - y_test) ** 2).mean()
results.append({"fold": fold, "ic": round(ic, 4), "mse": round(mse, 4)})
df = pd.DataFrame(results)
df.loc["mean"] = df.mean()
return df
@property
def n_estimators_used(self) -> int | None:
"""实际使用的树数量。"""
if self._model is None:
return None
return self._model.tree_count_
+142
View File
@@ -0,0 +1,142 @@
"""
特征工程因子 特征矩阵 + 目标标签
严禁使用未来数据所有变换基于 expanding window 或训练集统计
"""
import numpy as np
import pandas as pd
from sklearn.preprocessing import RobustScaler
class FeatureEngine:
"""
特征工程引擎
参数:
lookahead: 预测未来 N 个交易日
label_type: 'regression' | 'classification'
winsorize_pct: 去极值的分位数边界 (0.01, 0.99)
nan_threshold: NaN 占比超过此值的因子直接剔除
"""
def __init__(
self,
lookahead: int = 5,
label_type: str = "regression",
winsorize_pct: tuple[float, float] = (0.01, 0.99),
nan_threshold: float = 0.3,
):
self.lookahead = lookahead
self.label_type = label_type
self.winsorize_pct = winsorize_pct
self.nan_threshold = nan_threshold
self._scaler = RobustScaler()
self._scaler_fitted = False
self._valid_features: list[str] = []
# ── 标签构建 ──────────────────────────────────────────
def build_labels(self, price_df: pd.DataFrame) -> pd.Series:
"""
构建目标标签
regression: (close_{t+N} - close_t) / close_t * 100
classification: 1 if return > 0 else 0
"""
close = price_df["close"]
future = close.shift(-self.lookahead)
ret = (future - close) / close * 100
if self.label_type == "classification":
return (ret > 0).astype(int)
return ret.rename(f"y_fwd_{self.lookahead}")
# ── 特征构建 ──────────────────────────────────────────
def build(
self,
factor_df: pd.DataFrame,
price_df: pd.DataFrame,
fit: bool = True,
) -> tuple[pd.DataFrame, pd.Series]:
"""
构建特征矩阵 X 和标签 y
参数:
factor_df: 因子 DataFrame, index=trade_date, columns=因子名
price_df: 价格 DataFrame, 需有 'close'
fit: True=训练模式fit scaler + 记录有效特征False=预测模式
返回:
X, yy predict 模式下为 None
"""
X = factor_df.copy()
# 1. 剔除 NaN 率过高的列
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
# 2. 缺失值填充:前值填充 → 截面中位数
X = X.ffill().fillna(X.median())
# 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. 构建标签
y = self.build_labels(price_df) if fit else None
# 6. 对齐(删掉无法构建标签的行)
if fit:
valid_idx = X.index.intersection(y.dropna().index)
X = X.loc[valid_idx]
y = y.loc[valid_idx]
return X, y
# ── 多股票构建 ────────────────────────────────────────
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 = [], []
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:
continue
X, y = self.build(f_df, p_df, fit=True)
if X.empty:
continue
X["_ts_code"] = ts_code
X_parts.append(X)
y_parts.append(y)
if not X_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
View File
+126
View File
@@ -0,0 +1,126 @@
"""
LightGBM 模型封装
"""
import numpy as np
import pandas as pd
import lightgbm as lgb
from models.base import BaseModel
from sklearn.model_selection import TimeSeriesSplit
_DEFAULT_PARAMS = {
"objective": "regression",
"metric": "rmse",
"boosting_type": "gbdt",
"num_leaves": 15,
"learning_rate": 0.03,
"feature_fraction": 0.7,
"bagging_fraction": 0.7,
"bagging_freq": 5,
"verbose": -1,
"n_estimators": 1000,
"random_state": 42,
"min_data_in_leaf": 20,
}
class LightGBMModel(BaseModel):
"""LightGBM 回归模型。"""
name = "lightgbm"
def __init__(
self,
params: dict | None = None,
early_stopping: int = 50,
eval_ratio: float = 0.2,
random_seed: int = 42,
):
self.params = params or _DEFAULT_PARAMS.copy()
self.params["random_state"] = random_seed
self.early_stopping = early_stopping
self.eval_ratio = eval_ratio
self._model: lgb.Booster | None = None
self._feature_names: list[str] = []
def fit(self, X: pd.DataFrame, y: pd.Series) -> "LightGBMModel":
self._feature_names = list(X.columns)
n = len(X)
val_size = int(n * self.eval_ratio)
callbacks = []
eval_set = None
X_train, y_train = X, y
# 验证集足够大时才启用早停(至少 50 条)
if val_size >= 50:
split_idx = n - val_size
X_train, X_val = X.iloc[:split_idx], X.iloc[split_idx:]
y_train, y_val = y.iloc[:split_idx], y.iloc[split_idx:]
eval_set = [(X_val, y_val)]
callbacks = [
lgb.early_stopping(stopping_rounds=self.early_stopping, verbose=False),
lgb.log_evaluation(0),
]
self._model = lgb.LGBMRegressor(**self.params)
self._model.fit(
X_train, y_train,
eval_set=eval_set,
callbacks=callbacks if callbacks else None,
)
return self
def predict(self, X: pd.DataFrame) -> pd.Series:
if self._model is None:
raise RuntimeError("模型尚未训练")
preds = self._model.predict(X[self._feature_names])
return pd.Series(preds, index=X.index, name="pred")
def get_feature_importance(self, importance_type: str = "gain") -> pd.DataFrame:
"""特征重要性。importance_type: 'gain' | 'split'"""
if self._model is None:
return pd.DataFrame()
imp = self._model.booster_.feature_importance(importance_type=importance_type)
names = self._model.booster_.feature_name()
df = pd.DataFrame({"feature": names, "importance": imp})
df["importance_pct"] = df["importance"] / df["importance"].sum() * 100
return df.sort_values("importance", ascending=False)
def cv_evaluate(
self, X: pd.DataFrame, y: pd.Series, n_folds: int = 5
) -> pd.DataFrame:
"""
时间序列交叉验证评估 shuffle
返回每折的 IC (相关系数) MSE
"""
tscv = TimeSeriesSplit(n_splits=n_folds)
results = []
for fold, (train_idx, test_idx) in enumerate(tscv.split(X)):
X_train, X_test = X.iloc[train_idx], X.iloc[test_idx]
y_train, y_test = y.iloc[train_idx], y.iloc[test_idx]
model = LightGBMModel(
params=self.params,
early_stopping=self.early_stopping,
eval_ratio=0.0, # 不使用内部验证,直接全量训练
)
model.fit(X_train, y_train)
preds = model.predict(X_test)
ic = preds.corr(y_test)
mse = ((preds - y_test) ** 2).mean()
results.append({"fold": fold, "ic": round(ic, 4), "mse": round(mse, 4)})
df = pd.DataFrame(results)
df.loc["mean"] = df.mean()
return df
@property
def n_estimators_used(self) -> int | None:
"""实际使用的树数量(早停后可能 < n_estimators)。"""
if self._model is None:
return None
return self._model.booster_.current_iteration()
View File
+242
View File
@@ -0,0 +1,242 @@
"""
Optuna 优化引擎
统一接口optimizer.optimize(strategy_class, space, price_df, factor_df) OptimizationResult
"""
import copy
import time
import numpy as np
import optuna
import pandas as pd
from backtest.base import BaseStrategy
from backtest.report import BacktestReport
from backtest.vectorbt.engine import VectorBTEngine
from optimizer.objectives import Objective
from optimizer.result import OptimizationResult, WalkForwardResult
from optimizer.space import SearchSpace
# 抑制 Optuna 日志
optuna.logging.set_verbosity(optuna.logging.WARNING)
class OptunaEngine:
"""
Optuna 优化引擎
"""
def __init__(self, bt_engine: VectorBTEngine | None = None):
self.bt_engine = bt_engine or VectorBTEngine()
def optimize(
self,
strategy_class: type[BaseStrategy],
search_space: SearchSpace,
price_df: pd.DataFrame,
factor_df: pd.DataFrame | None = None,
metric: str = "sharpe",
n_trials: int = 100,
direction: str = "maximize",
sampler: optuna.samplers.BaseSampler | None = None,
) -> OptimizationResult:
"""
参数寻优
参数:
strategy_class: 策略类
search_space: 搜索空间
price_df: 价格数据
factor_df: 因子数据
metric: 优化目标
n_trials: 试验次数
direction: 'maximize' | 'minimize'
sampler: Optuna 采样器默认 TPESampler
"""
if sampler is None:
sampler = optuna.samplers.TPESampler(seed=42)
study = optuna.create_study(
direction=direction,
sampler=sampler,
)
objective = Objective(
strategy_class=strategy_class,
search_space=search_space,
price_df=price_df,
factor_df=factor_df,
bt_engine=self.bt_engine,
metric=metric,
)
t0 = time.time()
study.optimize(objective, n_trials=n_trials, show_progress_bar=True)
elapsed = time.time() - t0
# 用最优参数跑一次完整回测
best_params = study.best_params
try:
best_strategy = strategy_class(**best_params)
except TypeError:
valid = {k: v for k, v in best_params.items()
if k in strategy_class.__init__.__code__.co_varnames}
best_strategy = strategy_class(**valid)
best_report = self.bt_engine.run(
best_strategy,
price_df,
price_df if factor_df is None else factor_df,
)
# 参数重要性
try:
importance = optuna.importance.get_param_importances(study)
except Exception:
importance = {}
# 试验记录
trials_df = study.trials_dataframe()
return OptimizationResult(
best_params=study.best_params,
best_value=study.best_value,
metric=metric,
best_report=best_report,
trials_df=trials_df,
param_importance=importance,
)
def optimize_walk_forward(
self,
strategy_class: type[BaseStrategy],
search_space: SearchSpace,
price_df: pd.DataFrame,
factor_df: pd.DataFrame | None = None,
metric: str = "sharpe",
n_trials: int = 100,
train_window: int = 252 * 3,
test_window: int = 252,
) -> WalkForwardResult:
"""
滚动窗口优化Walk-Forward Analysis
每一步train_window 训练 test_window 验证 滑动
"""
if factor_df is None:
factor_df = price_df
n_total = len(price_df)
windows = []
test_equities = []
param_history = []
start = 0
while start + train_window + test_window <= n_total:
train_slice = slice(start, start + train_window)
test_slice = slice(start + train_window, start + train_window + test_window)
train_price = price_df.iloc[train_slice]
train_factor = factor_df.iloc[train_slice]
test_price = price_df.iloc[test_slice]
test_factor = factor_df.iloc[test_slice]
# 训练集上优化
opt_result = self.optimize(
strategy_class=strategy_class,
search_space=search_space,
price_df=train_price,
factor_df=train_factor,
metric=metric,
n_trials=n_trials,
)
# 测试集上验证
try:
test_strategy = strategy_class(**opt_result.best_params)
except TypeError:
valid = {k: v for k, v in opt_result.best_params.items()
if k in strategy_class.__init__.__code__.co_varnames}
test_strategy = strategy_class(**valid)
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)
train_idx = train_price.index
test_idx = test_price.index
windows.append({
"train_start": train_idx[0] if len(train_idx) > 0 else "",
"train_end": train_idx[-1] if len(train_idx) > 0 else "",
"test_start": test_idx[0] if len(test_idx) > 0 else "",
"test_end": test_idx[-1] if len(test_idx) > 0 else "",
"best_params": opt_result.best_params,
"best_value": opt_result.best_value,
"test_return": test_report.total_return,
"test_sharpe": test_report.sharpe_ratio,
"test_mdd": test_report.max_drawdown,
})
param_history.append(opt_result.best_params)
start += test_window
# 合并测试期权益曲线
consolidated = _merge_test_periods(test_equities, self.bt_engine.initial_capital)
# 参数稳定性
param_df = pd.DataFrame(param_history) if param_history else pd.DataFrame()
if not param_df.empty:
param_df.index.name = "window"
return WalkForwardResult(
windows=windows,
consolidated_report=consolidated,
param_stability=param_df,
)
def _merge_test_periods(
equity_list: list[pd.Series],
initial_capital: float = 100_000,
) -> BacktestReport | None:
"""拼接各窗口测试期权益曲线为一个连续序列。"""
if not equity_list:
return None
merged = pd.concat(equity_list)
merged = merged.sort_index()
merged = merged[~merged.index.duplicated()]
# 确保 DatetimeIndex
if not isinstance(merged.index, pd.DatetimeIndex):
merged.index = pd.to_datetime(merged.index, format="%Y%m%d")
dd = merged / merged.cummax() - 1
daily_ret = merged.pct_change().dropna()
years = max(len(daily_ret) / 252, 0.02)
total_ret = (merged.iloc[-1] / merged.iloc[0] - 1) * 100
cagr = ((total_ret / 100 + 1) ** (1 / years) - 1) * 100
mdd = dd.min() * 100
std_ret = daily_ret.std() * np.sqrt(252)
sharpe = (daily_ret.mean() * 252) / std_ret if std_ret > 0 else 0
calmar = cagr / abs(mdd) if abs(mdd) > 0 else 0
try:
monthly = merged.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),
equity_curve=merged,
drawdown_curve=dd,
monthly_returns=monthly,
)
+74
View File
@@ -0,0 +1,74 @@
"""
Optuna 目标函数
将策略实例化 回测 提取指标包装为 Optuna objective
"""
import optuna
import pandas as pd
from backtest.base import BaseStrategy
from backtest.vectorbt.engine import VectorBTEngine
from optimizer.space import SearchSpace
# 指标提取器:从 BacktestReport 取对应字段
_METRIC_EXTRACTORS = {
"sharpe": lambda r: r.sharpe_ratio,
"cagr": lambda r: r.cagr,
"calmar": lambda r: r.calmar_ratio,
"total_return": lambda r: r.total_return,
"return_over_dd": lambda r: abs(r.total_return / r.max_drawdown) if r.max_drawdown != 0 else 0.0,
"win_rate": lambda r: r.win_rate,
"profit_factor": lambda r: r.profit_factor,
}
class Objective:
"""
Optuna 目标函数可调用
用法:
obj = Objective(SMACrossStrategy, sma_cross_space, price_df, factor_df, metric="sharpe")
study = optuna.create_study(direction="maximize")
study.optimize(obj, n_trials=100)
"""
def __init__(
self,
strategy_class: type[BaseStrategy],
search_space: SearchSpace,
price_df: pd.DataFrame,
factor_df: pd.DataFrame | None = None,
bt_engine: VectorBTEngine | None = None,
metric: str = "sharpe",
):
self.strategy_class = strategy_class
self.search_space = search_space
self.price_df = price_df
self.factor_df = factor_df if factor_df is not None else price_df
self.bt_engine = bt_engine or VectorBTEngine()
self.metric = metric
self._extractor = _METRIC_EXTRACTORS.get(metric)
if self._extractor is None:
raise ValueError(f"不支持的指标: '{metric}'。可选: {list(_METRIC_EXTRACTORS)}")
def __call__(self, trial: optuna.Trial) -> float:
params = self.search_space.suggest(trial)
try:
strategy = self.strategy_class(**params)
except TypeError:
# 过滤不匹配的参数
valid = {k: v for k, v in params.items()
if k in self.strategy_class.__init__.__code__.co_varnames}
strategy = self.strategy_class(**valid)
report = self.bt_engine.run(strategy, self.price_df, self.factor_df)
value = self._extractor(report) # type: ignore
# 无效值处理
if value is None or (isinstance(value, float) and (pd.isna(value) or value == float("inf"))):
return float("-inf")
return float(value)
+86
View File
@@ -0,0 +1,86 @@
"""
策略优化快捷函数
为常用策略提供一键优化入口
"""
import pandas as pd
from backtest.vectorbt.engine import VectorBTEngine
from optimizer.engine import OptunaEngine
from optimizer.result import OptimizationResult, WalkForwardResult
from optimizer.space import (
sma_cross_space,
rsi_revert_space,
momentum_breakout_space,
factor_cross_space,
)
_DEFAULT_TRIALS = 100
def optimize_sma_cross(
price_df: pd.DataFrame,
factor_df: pd.DataFrame | None = None,
bt_engine: VectorBTEngine | None = None,
n_trials: int = _DEFAULT_TRIALS,
metric: str = "sharpe",
) -> OptimizationResult:
"""均线交叉策略参数寻优。"""
from backtest.strategies.sma_cross import SMACrossStrategy
return OptunaEngine(bt_engine).optimize(
SMACrossStrategy, sma_cross_space, price_df, factor_df, metric, n_trials,
)
def optimize_rsi_revert(
price_df: pd.DataFrame,
factor_df: pd.DataFrame | None = None,
bt_engine: VectorBTEngine | None = None,
n_trials: int = _DEFAULT_TRIALS,
metric: str = "sharpe",
) -> OptimizationResult:
"""RSI 反转策略参数寻优。"""
from backtest.strategies.rsi_mean_revert import RSIMeanRevertStrategy
return OptunaEngine(bt_engine).optimize(
RSIMeanRevertStrategy, rsi_revert_space, price_df, factor_df, metric, n_trials,
)
def optimize_momentum_breakout(
price_df: pd.DataFrame,
factor_df: pd.DataFrame | None = None,
bt_engine: VectorBTEngine | None = None,
n_trials: int = _DEFAULT_TRIALS,
metric: str = "sharpe",
) -> OptimizationResult:
"""动量突破策略参数寻优。"""
from backtest.strategies.momentum_breakout import MomentumBreakoutStrategy
return OptunaEngine(bt_engine).optimize(
MomentumBreakoutStrategy, momentum_breakout_space, price_df, factor_df, metric, n_trials,
)
def optimize_factor_cross(
price_df: pd.DataFrame,
factor_column: str,
factor_df: pd.DataFrame | None = None,
bt_engine: VectorBTEngine | None = None,
n_trials: int = _DEFAULT_TRIALS,
metric: str = "sharpe",
) -> OptimizationResult:
"""因子阈值交叉策略参数寻优。
参数:
factor_column: 因子列名 'momentum_20'
其余同 optimize_* 系列
"""
from backtest.strategies.factor_cross import FactorCrossStrategy
class _FCS(FactorCrossStrategy):
def __init__(self, buy_threshold=0, sell_threshold=None):
super().__init__(factor_column, buy_threshold, sell_threshold)
return OptunaEngine(bt_engine).optimize(
_FCS, factor_cross_space, price_df, factor_df, metric, n_trials,
)
+57
View File
@@ -0,0 +1,57 @@
"""
优化结果数据结构
"""
from dataclasses import dataclass, field
import pandas as pd
from backtest.report import BacktestReport
@dataclass
class OptimizationResult:
"""单次参数优化结果。"""
best_params: dict = field(default_factory=dict)
best_value: float = 0.0
metric: str = "sharpe"
best_report: BacktestReport | None = None
trials_df: pd.DataFrame = field(default_factory=pd.DataFrame)
param_importance: dict = field(default_factory=dict)
def summary(self) -> str:
lines = [
f"最优参数: {self.best_params}",
f"最优目标 ({self.metric}): {self.best_value:.4f}",
]
if self.best_report is not None:
lines.append(f"回测: {self.best_report.summary()}")
return "\n".join(lines)
@dataclass
class WalkForwardResult:
"""滚动窗口优化结果。"""
windows: list[dict] = field(default_factory=list)
consolidated_report: BacktestReport | None = None
param_stability: pd.DataFrame = field(default_factory=pd.DataFrame)
def summary(self) -> str:
n = len(self.windows)
lines = [f"Walk-Forward: {n} 个窗口"]
for w in self.windows:
lines.append(
f" {w['train_start']}~{w['train_end']}"
f"{w['test_start']}~{w['test_end']}"
f" | 参数={w.get('best_params', {})}"
f" | 收益={w.get('test_return', 0):.1f}%"
)
if self.consolidated_report is not None:
lines.append(f"整体: {self.consolidated_report.summary()}")
if not self.param_stability.empty:
stds = self.param_stability.std()
lines.append(f"参数稳定性(std): {dict(stds.round(2))}")
return "\n".join(lines)
+67
View File
@@ -0,0 +1,67 @@
"""
参数搜索空间定义
"""
from dataclasses import dataclass, field
import optuna
@dataclass
class SearchSpace:
"""参数搜索空间。"""
params: list[dict] = field(default_factory=list)
# 每个元素: {"name": str, "type": "int"|"float"|"categorical",
# "low": float, "high": float, "step": float, "choices": list}
def suggest(self, trial: optuna.Trial) -> dict:
"""从 trial 中采样一组参数。"""
result = {}
for p in self.params:
name = p["name"]
kind = p["type"]
if kind == "int":
low = p.get("low", 0)
high = p.get("high", 100)
step = p.get("step", 1)
result[name] = trial.suggest_int(name, int(low), int(high), step=int(step))
elif kind == "float":
low = p.get("low", 0.0)
high = p.get("high", 1.0)
result[name] = trial.suggest_float(name, float(low), float(high))
elif kind == "categorical":
choices = p.get("choices", [])
result[name] = trial.suggest_categorical(name, choices)
return result
# ── 预置搜索空间 ──────────────────────────────────────────
sma_cross_space = SearchSpace(params=[
{"name": "fast", "type": "int", "low": 2, "high": 30, "step": 1},
{"name": "slow", "type": "int", "low": 15, "high": 120, "step": 5},
])
rsi_revert_space = SearchSpace(params=[
{"name": "oversold", "type": "int", "low": 10, "high": 45, "step": 1},
{"name": "overbought", "type": "int", "low": 55, "high": 90, "step": 1},
])
momentum_breakout_space = SearchSpace(params=[
{"name": "lookback", "type": "int", "low": 10, "high": 60, "step": 5},
{"name": "exit_period", "type": "int", "low": 5, "high": 30, "step": 1},
])
factor_cross_space = SearchSpace(params=[
{"name": "buy_threshold", "type": "float", "low": -10.0, "high": 10.0},
{"name": "sell_threshold", "type": "float", "low": -10.0, "high": 10.0},
])
# 名称 → 空间映射
SPACES = {
"sma_cross": sma_cross_space,
"rsi_revert": rsi_revert_space,
"momentum_breakout": momentum_breakout_space,
"factor_cross": factor_cross_space,
}
View File
View File
+121
View File
@@ -0,0 +1,121 @@
"""
报告持久化模块
CLI 脚本输出的 Markdown 报告存入 DBmac_report
同一日期+研究对象的新报告入库时旧报告自动标记为失效is_active=0
"""
from datetime import datetime
from sqlalchemy import text
from database.connection import get_engine
from database.models import Report, create_all_tables
def save_report(
content: str,
title: str,
report_date: str | None = None,
subject_type: str = "daily",
subject_code: str = "",
) -> int:
"""
保存报告到 DB
如果同一天同一研究对象已有报告先将其标记为 is_active=0
然后插入新报告
参数:
content: Markdown 报告内容
title: 报告标题
report_date: 报告日期 YYYYMMDD默认今天
subject_type: 研究对象类型 (stock/index/sector/portfolio/daily)
subject_code: 研究对象代码 ( 000001.SZ 000300.SH)
返回:
新报告的 id
"""
report_date = report_date or datetime.now().strftime("%Y%m%d")
engine = get_engine()
# 确保表存在
create_all_tables()
# 将同日期+同对象的旧报告标记失效
with engine.connect() as conn:
conn.execute(
text(
"UPDATE {} SET is_active = 0 "
"WHERE report_date = :d AND subject_type = :st AND subject_code = :sc AND is_active = 1"
.format(Report.__tablename__)
),
{"d": report_date, "st": subject_type, "sc": subject_code},
)
conn.commit()
# 插入新报告
report = Report(
report_date=report_date,
title=title,
subject_type=subject_type,
subject_code=subject_code,
content=content,
created_at=datetime.now(),
is_active=1.0,
)
from sqlalchemy.orm import Session
with Session(engine) as session:
session.add(report)
session.commit()
report_id = report.id
session.expunge_all()
return report_id
def query_reports(
report_date: str | None = None,
subject_type: str | None = None,
subject_code: str | None = None,
active_only: bool = True,
limit: int = 20,
) -> list[dict]:
"""查询报告列表。"""
engine = get_engine()
table = Report.__tablename__
sql = "SELECT * FROM {} WHERE 1=1".format(table)
params = {}
if report_date:
sql += " AND report_date = :d"
params["d"] = report_date
if subject_type:
sql += " AND subject_type = :st"
params["st"] = subject_type
if subject_code:
sql += " AND subject_code = :sc"
params["sc"] = subject_code
if active_only:
sql += " AND is_active = 1"
sql += " ORDER BY id DESC LIMIT :lim"
params["lim"] = limit
with engine.connect() as conn:
rows = conn.execute(text(sql), params).fetchall()
result = []
for r in rows:
d = dict(r._mapping)
# DATE/DATETIME 列 → 字符串
for key in ("report_date", "created_at"):
val = d.get(key)
if hasattr(val, "strftime"):
fmt = "%Y%m%d" if key == "report_date" else "%Y-%m-%d %H:%M:%S"
d[key] = val.strftime(fmt)
result.append(d)
return result
+10
View File
@@ -0,0 +1,10 @@
akshare>=1.14.0
pandas>=2.0.0
sqlalchemy>=2.0
pymysql>=1.1.0
python-dotenv>=1.0
vectorbt>=0.5.0
optuna>=3.0
lightgbm>=4.0
catboost>=1.2
quantstats>=0.0.62
View File
View File