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

313 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
情绪因子计算引擎。
全链路:数据获取 → 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. 计算因子(对齐注册表,产出全部 4 个注册情绪因子,含 news_sent_20)
factor_dfs = {}
if not sentiment_df.empty:
for factor_cls, kwargs in [
(NewsSentimentFactor, {"window": 5, "sentiment_df": sentiment_df}),
(NewsSentimentFactor, {"window": 20, "sentiment_df": sentiment_df}),
(SentimentConfidenceFactor, {"window": 5, "sentiment_df": sentiment_df}),
(SentimentMomentumFactor, {"period": 5, "sentiment_df": sentiment_df}),
]:
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