Files
quanxiel/alpha/data_loader.py

283 lines
12 KiB
Python

"""
数据加载辅助模块 — 从量化数据库加载回测所需数据
支持的数据源:
- 日线行情 (daily + adj_factor → 后复权价格)
- 每日指标 (daily_basic → PE / PB / 市值 / 换手率)
- 财务指标 (fina_indicator → ROE / ROA / 毛利率)
- 指数行情 (index_daily → 基准收益)
若数据库连接失败,自动回退到随机模拟数据,便于 Notebook 演示完整研究流程。
用法:
from data_loader import DataLoader
loader = DataLoader()
price_data = loader.load_prices('2020-01-01', '2024-12-31')
pe_data = loader.load_factor('pe', '2020-01-01', '2024-12-31')
roe_data = loader.load_factor('roe', '2020-01-01', '2024-12-31')
bench = loader.load_benchmark('000300.SH', '2020-01-01', '2024-12-31')
"""
import sys
from pathlib import Path
from typing import Dict, List, Optional
import numpy as np
import pandas as pd
# 项目根目录 (quanxiel/)
_PROJECT_ROOT = Path(__file__).resolve().parent.parent
class DataLoader:
"""量化数据库数据加载器(数据库不可用时自动回退到模拟数据)"""
# 默认股票池(沪深300中代表性的蓝筹成分股)
DEFAULT_STOCKS: List[str] = [
# 银行 / 非银
"600000.SH", "600015.SH", "600016.SH", "600036.SH", "601166.SH",
"601288.SH", "601398.SH", "601318.SH", "601601.SH", "601628.SH",
# 消费 / 医药
"600519.SH", "000858.SZ", "000651.SZ", "000333.SZ", "600887.SH",
"600276.SH", "601888.SH",
# 工业 / 制造
"600031.SH", "600048.SH", "600104.SH", "600309.SH", "600585.SH",
"601668.SH", "601088.SH", "600028.SH", "601857.SH",
# 科技 / 成长
"000063.SZ", "002415.SZ", "300750.SZ", "600009.SH", "600050.SH",
]
# daily_basic 表可用字段
DAILY_BASIC_FIELDS = {
"pe", "pe_ttm", "pb", "ps", "ps_ttm", "dv_ratio", "dv_ttm",
"total_mv", "circ_mv", "turnover_rate", "volume_ratio",
}
# fina_indicator 表可用字段
FINA_FIELDS = {
"roe", "roe_dt", "roa", "roa2", "roic", "gross_margin",
"eps", "dt_eps", "current_ratio", "quick_ratio", "debt_to_assets",
"assets_turn", "cfo_to_or", "ocfps",
}
def __init__(self, stocks: Optional[List[str]] = None, seed: int = 42):
self.stocks = stocks or self.DEFAULT_STOCKS
self.seed = seed
self._conn = None
self._db_ok: Optional[bool] = None
# ==================================================================
# 公开接口
# ==================================================================
def load_prices(self, start_date: str, end_date: str) -> pd.DataFrame:
"""
加载后复权收盘价。
返回 DataFrame: index=date, columns=ts_code
"""
if self._ensure_db():
try:
return self._load_prices_db(start_date, end_date)
except Exception as e:
print(f"[DataLoader] 价格数据加载失败,回退到模拟数据: {e}")
return self._simulate_prices(start_date, end_date)
def load_factor(self, name: str, start_date: str, end_date: str) -> pd.DataFrame:
"""
加载指定因子/指标截面数据。
返回 DataFrame: index=date, columns=ts_code
"""
if name in self.DAILY_BASIC_FIELDS and self._ensure_db():
try:
return self._load_daily_basic_field(name, start_date, end_date)
except Exception as e:
print(f"[DataLoader] 指标 {name} 加载失败,回退到模拟数据: {e}")
elif name in self.FINA_FIELDS and self._ensure_db():
try:
return self._load_fina_field(name, start_date, end_date)
except Exception as e:
print(f"[DataLoader] 财务指标 {name} 加载失败,回退到模拟数据: {e}")
return self._simulate_factor(name, start_date, end_date)
def load_factors(self, names: List[str], start_date: str, end_date: str) -> Dict[str, pd.DataFrame]:
"""批量加载多个因子,返回 {name: DataFrame}"""
return {n: self.load_factor(n, start_date, end_date) for n in names}
def load_benchmark(
self, code: str = "000300.SH", start_date: str = "", end_date: str = ""
) -> pd.Series:
"""
加载基准指数收盘价。
返回 Series: index=date, values=close
"""
if start_date and end_date and self._ensure_db():
try:
return self._load_index_db(code, start_date, end_date)
except Exception as e:
print(f"[DataLoader] 指数 {code} 加载失败,回退到模拟数据: {e}")
# 模拟基准:用股票池均值价格代替
prices = self._simulate_prices(start_date or "2020-01-01", end_date or "2024-12-31")
return prices.mean(axis=1).rename(code)
def get_trade_dates(self, start_date: str, end_date: str) -> pd.DatetimeIndex:
"""获取交易日列表(数据库不可用时使用工作日)"""
if self._ensure_db():
try:
conn = self._get_conn()
query = """
SELECT DISTINCT cal_date FROM trade_cal
WHERE is_open = 1 AND cal_date BETWEEN %s AND %s
ORDER BY cal_date
"""
dates = pd.read_sql(query, conn, params=[start_date, end_date])
return pd.to_datetime(dates["cal_date"])
except Exception:
pass
return pd.bdate_range(start_date, end_date)
# ==================================================================
# 数据库连接
# ==================================================================
def _get_conn(self):
if self._conn is None:
sys.path.insert(0, str(_PROJECT_ROOT / "quantitative_data"))
from config import DB_CONFIG
import psycopg2
self._conn = psycopg2.connect(
host=DB_CONFIG["host"],
port=DB_CONFIG["port"],
database=DB_CONFIG["database"],
user=DB_CONFIG["user"],
password=DB_CONFIG["password"],
connect_timeout=5,
)
return self._conn
def _ensure_db(self) -> bool:
"""检测数据库是否可用(结果缓存)"""
if self._db_ok is not None:
return self._db_ok
try:
self._get_conn()
self._db_ok = True
except Exception:
self._conn = None
self._db_ok = False
print("[DataLoader] 数据库连接失败,使用模拟数据。"
"请检查 quantitative_data/.env 配置。")
return self._db_ok
# ==================================================================
# 数据库加载实现
# ==================================================================
def _stock_sql_placeholders(self) -> str:
return ",".join(["%s"] * len(self.stocks))
def _load_prices_db(self, start_date: str, end_date: str) -> pd.DataFrame:
conn = self._get_conn()
placeholders = self._stock_sql_placeholders()
query = f"""
SELECT d.ts_code, d.trade_date,
d.close * a.adj_factor AS adj_close
FROM daily d
JOIN adj_factor a
ON d.ts_code = a.ts_code AND d.trade_date = a.trade_date
WHERE d.ts_code IN ({placeholders})
AND d.trade_date BETWEEN %s AND %s
ORDER BY d.trade_date, d.ts_code
"""
df = pd.read_sql(query, conn, params=[*self.stocks, start_date, end_date])
if df.empty:
raise ValueError("查询无数据")
pivot = df.pivot(index="trade_date", columns="ts_code", values="adj_close")
pivot.index = pd.to_datetime(pivot.index)
pivot = pivot.sort_index()
# 只保留有数据的列
return pivot.dropna(how="all", axis=1)
def _load_daily_basic_field(self, field: str, start_date: str, end_date: str) -> pd.DataFrame:
conn = self._get_conn()
placeholders = self._stock_sql_placeholders()
query = f"""
SELECT ts_code, trade_date, {field}
FROM daily_basic
WHERE ts_code IN ({placeholders})
AND trade_date BETWEEN %s AND %s
ORDER BY trade_date, ts_code
"""
df = pd.read_sql(query, conn, params=[*self.stocks, start_date, end_date])
if df.empty:
raise ValueError(f"查询 {field} 无数据")
pivot = df.pivot(index="trade_date", columns="ts_code", values=field)
pivot.index = pd.to_datetime(pivot.index)
pivot = pivot.sort_index()
return pivot.replace([np.inf, -np.inf], np.nan)
def _load_fina_field(self, field: str, start_date: str, end_date: str) -> pd.DataFrame:
conn = self._get_conn()
placeholders = self._stock_sql_placeholders()
query = f"""
SELECT ts_code, end_date, {field}
FROM fina_indicator
WHERE ts_code IN ({placeholders})
AND end_date BETWEEN %s AND %s
ORDER BY end_date, ts_code
"""
df = pd.read_sql(query, conn, params=[*self.stocks, start_date, end_date])
if df.empty:
raise ValueError(f"查询财务指标 {field} 无数据")
pivot = df.pivot(index="end_date", columns="ts_code", values=field)
pivot.index = pd.to_datetime(pivot.index)
# 重采样到交易日并向前填充
trade_dates = self.get_trade_dates(start_date, end_date)
pivot = pivot.sort_index().reindex(trade_dates).ffill()
return pivot.replace([np.inf, -np.inf], np.nan)
def _load_index_db(self, code: str, start_date: str, end_date: str) -> pd.Series:
conn = self._get_conn()
query = """
SELECT trade_date, close
FROM index_daily
WHERE ts_code = %s AND trade_date BETWEEN %s AND %s
ORDER BY trade_date
"""
df = pd.read_sql(query, conn, params=[code, start_date, end_date])
if df.empty:
raise ValueError(f"指数 {code} 无数据")
s = pd.Series(df["close"].values, index=pd.to_datetime(df["trade_date"]))
return s.sort_index().rename(code)
# ==================================================================
# 模拟数据(数据库不可用时的回退)
# ==================================================================
def _simulate_prices(self, start_date: str, end_date: str) -> pd.DataFrame:
np.random.seed(self.seed)
dates = pd.bdate_range(start_date, end_date)
n_dates, n_stocks = len(dates), len(self.stocks)
# 几何布朗运动模拟股价
rets = np.random.randn(n_dates, n_stocks) * 0.02
rets[:, 0] *= 0.5 # 让第一只股票波动小(制造多样化)
rets[:, 1] *= 1.5 # 让第二只股票波动大
prices = 100 * np.exp(np.cumsum(rets, axis=0))
df = pd.DataFrame(prices, index=dates, columns=self.stocks)
return df
def _simulate_factor(self, name: str, start_date: str, end_date: str) -> pd.DataFrame:
np.random.seed(self.seed)
dates = pd.bdate_range(start_date, end_date)
n = len(self.stocks)
data = np.random.randn(len(dates), n)
# 部分指标需要为正(估值/基本面)
positive_fields = {
"pe", "pe_ttm", "pb", "ps", "ps_ttm", "dv_ratio", "dv_ttm",
"total_mv", "circ_mv", "turnover_rate", "volume_ratio",
"roe", "roe_dt", "roa", "roa2", "roic", "gross_margin",
"eps", "dt_eps", "current_ratio", "quick_ratio", "debt_to_assets",
"assets_turn", "cfo_to_or", "ocfps",
}
if name in positive_fields:
data = np.abs(data) + 0.5
return pd.DataFrame(data, index=dates, columns=self.stocks)