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:
@@ -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: 日线 DataFrame,index 为 trade_date,
|
||||
列至少包含 OHLCV 等基础字段。
|
||||
|
||||
返回:
|
||||
pd.Series,index 与 df 对齐,值为因子值。
|
||||
"""
|
||||
...
|
||||
|
||||
def get_required_columns(self) -> list[str]:
|
||||
"""返回计算所需的列名列表。子类可覆盖。"""
|
||||
return ["close"]
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(name='{self.name}')"
|
||||
@@ -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: 因子实例列表
|
||||
|
||||
返回:
|
||||
DataFrame,index=trade_date,columns=因子名
|
||||
"""
|
||||
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 表示全部
|
||||
|
||||
返回:
|
||||
DataFrame,index=ts_code,columns=因子名
|
||||
"""
|
||||
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
|
||||
@@ -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")
|
||||
@@ -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: 财务数据 DataFrame,columns 含 '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")
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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: 情绪 DataFrame(date 列 + 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")
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user