diff --git a/src/quant_engine/data_adapter.py b/src/quant_engine/data_adapter.py index 9181700..041b6c3 100644 --- a/src/quant_engine/data_adapter.py +++ b/src/quant_engine/data_adapter.py @@ -287,6 +287,8 @@ def prepare_execution_inputs( df: pd.DataFrame, stock_col: str = "stock_code", date_col: str = "trade_date", + *, + price_col: str = "close", ) -> tuple[pd.DataFrame, pd.DataFrame]: """长表行情 → execution 输入(prices + volumes 宽表)。 @@ -294,10 +296,11 @@ def prepare_execution_inputs( df: 长表行情(含 close / volume 列,Tushare rename 后) stock_col: 股票代码列名 date_col: 日期列名 + price_col: 执行价字段,默认 close;防前视研究可显式选择下一交易日 open Returns: (prices_wide, volumes_wide): - - prices_wide: date × stock_code,值=close + - prices_wide: date × stock_code,值=price_col - volumes_wide: date × stock_code,值=volume(若无 volume 列则全 1.0) Examples: @@ -313,9 +316,9 @@ def prepare_execution_inputs( """ if df.empty: return pd.DataFrame(), pd.DataFrame() - if "close" not in df.columns: - raise ValueError(f"prepare_execution_inputs: 缺 close 列,实际列={list(df.columns)}") - prices = long_to_wide(df, value_col="close", date_col=date_col, stock_col=stock_col) + if price_col not in df.columns: + raise ValueError(f"prepare_execution_inputs: 缺 {price_col} 列,实际列={list(df.columns)}") + prices = long_to_wide(df, value_col=price_col, date_col=date_col, stock_col=stock_col) if "volume" in df.columns: volumes = long_to_wide(df, value_col="volume", date_col=date_col, stock_col=stock_col) else: diff --git a/src/quant_engine/research_pipeline.py b/src/quant_engine/research_pipeline.py new file mode 100644 index 0000000..0d96496 --- /dev/null +++ b/src/quant_engine/research_pipeline.py @@ -0,0 +1,217 @@ +"""可信研究链路:因子分数经交易日历滞后后进入执行审计。 + +本模块只编排现有组合构建与执行组件,不连接账户、券商或实盘订单。 +时间契约借鉴 Qlib 的 prediction/trade time 分离与 Backtrader 的 next-bar +执行语义:signal_date 上形成的目标权重,默认最早在下一交易时点执行。 +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np +import pandas as pd +from pandas.api.types import is_numeric_dtype + +from quant_engine.execution import ( + ExecutionConfig, + ExecutionSimulationResult, + simulate_multi_day_with_audit, +) +from quant_engine.portfolio_construction import scores_to_weight_table + +__all__ = [ + "TargetWeightSchedule", + "FactorExecutionResult", + "schedule_target_weights", + "run_factor_execution_research", +] + + +@dataclass(frozen=True, slots=True, eq=False) +class TargetWeightSchedule: + """保留决策时间和执行时间的目标权重调度快照。""" + + decision_weights: pd.DataFrame + signal_to_execution: pd.Series + execution_weights: pd.DataFrame + lag_sessions: int + + +@dataclass(frozen=True, slots=True, eq=False) +class FactorExecutionResult: + """因子到执行审计的一次可复现研究结果。""" + + factor_scores: pd.DataFrame + execution_prices: pd.DataFrame + schedule: TargetWeightSchedule + execution_price_field: str + execution: ExecutionSimulationResult + + +def _validate_datetime_index(index: pd.Index, name: str) -> pd.DatetimeIndex: + if not isinstance(index, pd.DatetimeIndex): + raise TypeError(f"{name} must use a DatetimeIndex") + if not index.is_unique: + raise ValueError(f"{name} must contain unique sessions") + if not index.is_monotonic_increasing: + raise ValueError(f"{name} must be in chronological order") + return index + + +def _validate_decision_weights(decision_weights: pd.DataFrame) -> None: + if not isinstance(decision_weights, pd.DataFrame): + raise TypeError( + f"decision_weights must be a pandas DataFrame, got {type(decision_weights).__name__}" + ) + _validate_datetime_index(decision_weights.index, "decision_weights index") + if not decision_weights.columns.is_unique: + raise ValueError("decision_weights must contain unique asset labels") + if not all(is_numeric_dtype(dtype) for dtype in decision_weights.dtypes): + raise TypeError("decision_weights must contain numeric values") + values = decision_weights.to_numpy(dtype=float) + if not np.isfinite(values).all() or (values < 0).any(): + raise ValueError("decision_weights must be finite and non-negative") + if (decision_weights.sum(axis=1) > 1.0 + 1e-12).any(): + raise ValueError("decision_weights rows must sum to at most 1.0") + + +def _validate_execution_prices(execution_prices: pd.DataFrame) -> pd.DatetimeIndex: + if not isinstance(execution_prices, pd.DataFrame): + raise TypeError( + f"execution_prices must be a pandas DataFrame, got {type(execution_prices).__name__}" + ) + calendar = _validate_datetime_index(execution_prices.index, "execution_prices index") + if not execution_prices.columns.is_unique: + raise ValueError("execution_prices must contain unique asset labels") + if not all(is_numeric_dtype(dtype) for dtype in execution_prices.dtypes): + raise TypeError("execution_prices must contain numeric values") + return calendar + + +def schedule_target_weights( + decision_weights: pd.DataFrame, + trading_calendar: pd.DatetimeIndex, + *, + lag_sessions: int = 1, +) -> TargetWeightSchedule: + """将信号日目标权重映射到后续真实交易日,不做整数行盲移位。 + + 所有信号日必须属于 ``trading_calendar``,且日历必须包含每个信号对应的 + 未来执行日;无法执行的末尾信号会显式失败,避免被静默丢弃。 + """ + _validate_decision_weights(decision_weights) + calendar = _validate_datetime_index(trading_calendar, "trading_calendar") + if isinstance(lag_sessions, bool) or not isinstance(lag_sessions, int) or lag_sessions <= 0: + raise ValueError("lag_sessions must be a positive integer") + + decision_snapshot = decision_weights.copy(deep=True) + if decision_snapshot.empty: + execution_weights = decision_snapshot.copy(deep=True) + execution_weights.index = pd.DatetimeIndex([], name="execution_date") + mapping = pd.Series( + calendar[:0], + index=decision_snapshot.index.copy(), + name="execution_date", + ) + return TargetWeightSchedule( + decision_weights=decision_snapshot, + signal_to_execution=mapping, + execution_weights=execution_weights, + lag_sessions=lag_sessions, + ) + + signal_positions = calendar.get_indexer(decision_snapshot.index) + if (signal_positions < 0).any(): + missing = decision_snapshot.index[signal_positions < 0] + raise ValueError( + "signal dates must be trading sessions; missing=" + + ", ".join(str(date) for date in missing) + ) + + execution_positions = signal_positions + lag_sessions + if (execution_positions >= len(calendar)).any(): + unavailable = decision_snapshot.index[execution_positions >= len(calendar)] + raise ValueError( + "trading_calendar lacks a future execution session for signal dates: " + + ", ".join(str(date) for date in unavailable) + ) + + execution_dates = calendar.take(execution_positions) + signal_to_execution = pd.Series( + execution_dates, + index=decision_snapshot.index.copy(), + name="execution_date", + ) + execution_weights = decision_snapshot.copy(deep=True) + execution_weights.index = pd.DatetimeIndex(execution_dates, name="execution_date") + return TargetWeightSchedule( + decision_weights=decision_snapshot, + signal_to_execution=signal_to_execution, + execution_weights=execution_weights, + lag_sessions=lag_sessions, + ) + + +def run_factor_execution_research( + factor_scores: pd.DataFrame, + execution_prices: pd.DataFrame, + *, + top_k: int, + execution_price_field: str, + lag_sessions: int = 1, + gross_exposure: float = 1.0, + largest: bool = True, + initial_cash: float = 1_000_000.0, + config: ExecutionConfig | None = None, +) -> FactorExecutionResult: + """运行因子分数 → 目标权重 → 下一交易时点 → 执行审计链路。 + + ``execution_prices`` 必须代表实际拟执行时点的价格矩阵,例如日频研究中 + signal 日收盘生成分数后使用下一交易日 ``open``。价格字段名称被保存在 + 结果元数据中,但函数不会猜测或重写价格语义。 + """ + price_field = execution_price_field.strip() + if not price_field: + raise ValueError("execution_price_field must be non-empty") + calendar = _validate_execution_prices(execution_prices) + + factor_snapshot = factor_scores.copy(deep=True) + decision_weights = scores_to_weight_table( + factor_snapshot, + top_k, + gross_exposure=gross_exposure, + largest=largest, + ) + schedule = schedule_target_weights( + decision_weights, + calendar, + lag_sessions=lag_sessions, + ) + price_snapshot = execution_prices.copy(deep=True) + + target_history: list[tuple[str, dict[str, float]]] = [] + price_history: list[tuple[str, dict[str, float]]] = [] + for execution_date, weights in schedule.execution_weights.iterrows(): + date_label = str(pd.Timestamp(execution_date)) + target_history.append( + (date_label, {asset: float(weight) for asset, weight in weights.items()}) + ) + prices = price_snapshot.loc[execution_date] + price_history.append( + (date_label, {asset: float(price) for asset, price in prices.items()}) + ) + + execution = simulate_multi_day_with_audit( + target_history, + price_history, + initial_cash, + config, + ) + return FactorExecutionResult( + factor_scores=factor_snapshot, + execution_prices=price_snapshot, + schedule=schedule, + execution_price_field=price_field, + execution=execution, + )