283 lines
12 KiB
Python
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) |