""" 数据加载辅助模块 — 从量化数据库加载回测所需数据 支持的数据源: - 日线行情 (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)