feat: schedule factor weights for next-session execution
This commit is contained in:
@@ -287,6 +287,8 @@ def prepare_execution_inputs(
|
|||||||
df: pd.DataFrame,
|
df: pd.DataFrame,
|
||||||
stock_col: str = "stock_code",
|
stock_col: str = "stock_code",
|
||||||
date_col: str = "trade_date",
|
date_col: str = "trade_date",
|
||||||
|
*,
|
||||||
|
price_col: str = "close",
|
||||||
) -> tuple[pd.DataFrame, pd.DataFrame]:
|
) -> tuple[pd.DataFrame, pd.DataFrame]:
|
||||||
"""长表行情 → execution 输入(prices + volumes 宽表)。
|
"""长表行情 → execution 输入(prices + volumes 宽表)。
|
||||||
|
|
||||||
@@ -294,10 +296,11 @@ def prepare_execution_inputs(
|
|||||||
df: 长表行情(含 close / volume 列,Tushare rename 后)
|
df: 长表行情(含 close / volume 列,Tushare rename 后)
|
||||||
stock_col: 股票代码列名
|
stock_col: 股票代码列名
|
||||||
date_col: 日期列名
|
date_col: 日期列名
|
||||||
|
price_col: 执行价字段,默认 close;防前视研究可显式选择下一交易日 open
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(prices_wide, volumes_wide):
|
(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)
|
- volumes_wide: date × stock_code,值=volume(若无 volume 列则全 1.0)
|
||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
@@ -313,9 +316,9 @@ def prepare_execution_inputs(
|
|||||||
"""
|
"""
|
||||||
if df.empty:
|
if df.empty:
|
||||||
return pd.DataFrame(), pd.DataFrame()
|
return pd.DataFrame(), pd.DataFrame()
|
||||||
if "close" not in df.columns:
|
if price_col not in df.columns:
|
||||||
raise ValueError(f"prepare_execution_inputs: 缺 close 列,实际列={list(df.columns)}")
|
raise ValueError(f"prepare_execution_inputs: 缺 {price_col} 列,实际列={list(df.columns)}")
|
||||||
prices = long_to_wide(df, value_col="close", date_col=date_col, stock_col=stock_col)
|
prices = long_to_wide(df, value_col=price_col, date_col=date_col, stock_col=stock_col)
|
||||||
if "volume" in df.columns:
|
if "volume" in df.columns:
|
||||||
volumes = long_to_wide(df, value_col="volume", date_col=date_col, stock_col=stock_col)
|
volumes = long_to_wide(df, value_col="volume", date_col=date_col, stock_col=stock_col)
|
||||||
else:
|
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