feat(quant): M7.3 研究行情口径显式化(默认不复权 none,可切 qfq)
- DailyBarRepository.get_range_many / stream_range_many_columns 增加 adjust 参数 (默认 'none')→ SQL 层过滤口径,消除 stock_daily 混 source/adjust 污染因子的风险 - ResearchSpec / SelectionQuery 增加 price_adjustment(none|qfq),随 config_snapshot 落库可溯源;ResearchService._load_daily 与 SelectionService 装配按口径取数 - tests/test_price_adjustment.py:repo 读取按 adjust 过滤(none/qfq 各自命中)、 spec 默认与字段记录;全量 pytest 通过
This commit is contained in:
@@ -62,6 +62,7 @@ class SelectionService:
|
||||
as_of - timedelta(days=query.warmup_days),
|
||||
as_of,
|
||||
columns,
|
||||
adjust=query.price_adjustment,
|
||||
)
|
||||
financial: dict[str, FinancialIndicator] = {}
|
||||
if query.method == "condition" and self._uses_fundamental(query):
|
||||
|
||||
@@ -62,6 +62,10 @@ class ResearchSpec(BaseModel):
|
||||
|
||||
type: str = Field(default="backtest", pattern="^(factor_test|backtest)$")
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
price_adjustment: str = Field(
|
||||
default="none", pattern="^(none|qfq)$",
|
||||
description="研究行情口径:none 不复权(默认)/ qfq 前复权(result 与 config_snapshot 中显式)",
|
||||
)
|
||||
factors: list[FactorSpec] = Field(min_length=1)
|
||||
selection: SelectionSpec = SelectionSpec()
|
||||
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
|
||||
|
||||
@@ -25,6 +25,10 @@ class SelectionQuery(BaseModel):
|
||||
"""一次选股查询(v2 §14.2 Selection 输入)。"""
|
||||
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
price_adjustment: str = Field(
|
||||
default="none", pattern="^(none|qfq)$",
|
||||
description="行情口径:none 不复权(默认)/ qfq 前复权(结果 config_snapshot 中显式)",
|
||||
)
|
||||
# 研究时点:None → 引擎用 <= 今天最近可用交易日;显式给历史日期即做历史选股
|
||||
as_of: date | None = Field(
|
||||
default=None, description="选股时点;历史回测/解释用具体日期,当前选股可留空"
|
||||
|
||||
@@ -42,8 +42,14 @@ class DailyBarRepository(Protocol):
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]: ...
|
||||
|
||||
def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]:
|
||||
"""批量区间查询(研究服务装配面板用,避免逐只查询)。"""
|
||||
def get_range_many(
|
||||
self,
|
||||
symbols: Sequence[str],
|
||||
start: date,
|
||||
end: date,
|
||||
adjust: str = "none",
|
||||
) -> list[DailyBar]:
|
||||
"""批量区间查询(研究装配面板用);adjust 指定行情口径(none 不复权 / qfq)。"""
|
||||
|
||||
def stream_range_many_columns(
|
||||
self,
|
||||
@@ -51,6 +57,7 @@ class DailyBarRepository(Protocol):
|
||||
start: date,
|
||||
end: date,
|
||||
columns: Sequence[str],
|
||||
adjust: str = "none",
|
||||
) -> Iterator[tuple]:
|
||||
"""流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。
|
||||
|
||||
|
||||
@@ -158,16 +158,24 @@ class SqlAlchemyDailyBarRepository:
|
||||
).all()
|
||||
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]:
|
||||
rows = self._session.scalars(
|
||||
def get_range_many(
|
||||
self,
|
||||
symbols: Sequence[str],
|
||||
start: date,
|
||||
end: date,
|
||||
adjust: str = "none",
|
||||
) -> list[DailyBar]:
|
||||
stmt = (
|
||||
select(StockDailyModel)
|
||||
.where(
|
||||
StockDailyModel.symbol.in_(list(symbols)),
|
||||
StockDailyModel.trade_date >= start,
|
||||
StockDailyModel.trade_date <= end,
|
||||
StockDailyModel.adjust == adjust, # 研究主口径:不复权(v2 §8)
|
||||
)
|
||||
.order_by(StockDailyModel.trade_date)
|
||||
).all()
|
||||
)
|
||||
rows = self._session.scalars(stmt).all()
|
||||
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
def stream_range_many_columns(
|
||||
@@ -176,6 +184,7 @@ class SqlAlchemyDailyBarRepository:
|
||||
start: date,
|
||||
end: date,
|
||||
columns: Sequence[str],
|
||||
adjust: str = "none",
|
||||
) -> Iterator[tuple]:
|
||||
"""流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。
|
||||
|
||||
@@ -193,6 +202,7 @@ class SqlAlchemyDailyBarRepository:
|
||||
StockDailyModel.symbol.in_(list(symbols)),
|
||||
StockDailyModel.trade_date >= start,
|
||||
StockDailyModel.trade_date <= end,
|
||||
StockDailyModel.adjust == adjust,
|
||||
)
|
||||
.order_by(StockDailyModel.symbol, StockDailyModel.trade_date)
|
||||
.execution_options(yield_per=20000)
|
||||
|
||||
@@ -75,6 +75,7 @@ def load_daily_df(
|
||||
start: date,
|
||||
end: date,
|
||||
columns: list[str],
|
||||
adjust: str = "none",
|
||||
) -> pd.DataFrame:
|
||||
"""从 Repository 装配行情长表(供研究/选股共用)。
|
||||
|
||||
@@ -86,12 +87,14 @@ def load_daily_df(
|
||||
streamer = getattr(daily_repo, "stream_range_many_columns", None)
|
||||
if streamer is not None:
|
||||
try:
|
||||
return _frame_from_stream(streamer(symbols, start, end, sorted(columns)), sorted(columns))
|
||||
return _frame_from_stream(
|
||||
streamer(symbols, start, end, sorted(columns), adjust=adjust), sorted(columns)
|
||||
)
|
||||
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
|
||||
pass
|
||||
get_many = getattr(daily_repo, "get_range_many", None)
|
||||
if get_many is not None:
|
||||
bars = list(get_many(symbols, start, end))
|
||||
bars = list(get_many(symbols, start, end, adjust=adjust))
|
||||
else: # 兜底:逐只查询
|
||||
bars = []
|
||||
for sym in symbols:
|
||||
@@ -134,5 +137,10 @@ class ResearchService:
|
||||
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
|
||||
required = self._engine.required_columns(spec)
|
||||
return load_daily_df(
|
||||
self._daily_repo, [s.symbol for s in stocks], data_start, end, sorted(required)
|
||||
self._daily_repo,
|
||||
[s.symbol for s in stocks],
|
||||
data_start,
|
||||
end,
|
||||
sorted(required),
|
||||
adjust=spec.price_adjustment,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user