- 新增 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 扩充配置项
313 lines
11 KiB
Python
313 lines
11 KiB
Python
"""
|
||
情绪因子计算引擎。
|
||
|
||
全链路:数据获取 → 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
|