feat: schedule factor weights for next-session execution
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user