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 - timedelta(days=query.warmup_days),
|
||||||
as_of,
|
as_of,
|
||||||
columns,
|
columns,
|
||||||
|
adjust=query.price_adjustment,
|
||||||
)
|
)
|
||||||
financial: dict[str, FinancialIndicator] = {}
|
financial: dict[str, FinancialIndicator] = {}
|
||||||
if query.method == "condition" and self._uses_fundamental(query):
|
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)$")
|
type: str = Field(default="backtest", pattern="^(factor_test|backtest)$")
|
||||||
universe: UniverseSpec = UniverseSpec()
|
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)
|
factors: list[FactorSpec] = Field(min_length=1)
|
||||||
selection: SelectionSpec = SelectionSpec()
|
selection: SelectionSpec = SelectionSpec()
|
||||||
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
|
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
|
||||||
|
|||||||
@@ -25,6 +25,10 @@ class SelectionQuery(BaseModel):
|
|||||||
"""一次选股查询(v2 §14.2 Selection 输入)。"""
|
"""一次选股查询(v2 §14.2 Selection 输入)。"""
|
||||||
|
|
||||||
universe: UniverseSpec = UniverseSpec()
|
universe: UniverseSpec = UniverseSpec()
|
||||||
|
price_adjustment: str = Field(
|
||||||
|
default="none", pattern="^(none|qfq)$",
|
||||||
|
description="行情口径:none 不复权(默认)/ qfq 前复权(结果 config_snapshot 中显式)",
|
||||||
|
)
|
||||||
# 研究时点:None → 引擎用 <= 今天最近可用交易日;显式给历史日期即做历史选股
|
# 研究时点:None → 引擎用 <= 今天最近可用交易日;显式给历史日期即做历史选股
|
||||||
as_of: date | None = Field(
|
as_of: date | None = Field(
|
||||||
default=None, description="选股时点;历史回测/解释用具体日期,当前选股可留空"
|
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(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(
|
def stream_range_many_columns(
|
||||||
self,
|
self,
|
||||||
@@ -51,6 +57,7 @@ class DailyBarRepository(Protocol):
|
|||||||
start: date,
|
start: date,
|
||||||
end: date,
|
end: date,
|
||||||
columns: Sequence[str],
|
columns: Sequence[str],
|
||||||
|
adjust: str = "none",
|
||||||
) -> Iterator[tuple]:
|
) -> Iterator[tuple]:
|
||||||
"""流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。
|
"""流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。
|
||||||
|
|
||||||
|
|||||||
@@ -158,16 +158,24 @@ class SqlAlchemyDailyBarRepository:
|
|||||||
).all()
|
).all()
|
||||||
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
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]:
|
def get_range_many(
|
||||||
rows = self._session.scalars(
|
self,
|
||||||
|
symbols: Sequence[str],
|
||||||
|
start: date,
|
||||||
|
end: date,
|
||||||
|
adjust: str = "none",
|
||||||
|
) -> list[DailyBar]:
|
||||||
|
stmt = (
|
||||||
select(StockDailyModel)
|
select(StockDailyModel)
|
||||||
.where(
|
.where(
|
||||||
StockDailyModel.symbol.in_(list(symbols)),
|
StockDailyModel.symbol.in_(list(symbols)),
|
||||||
StockDailyModel.trade_date >= start,
|
StockDailyModel.trade_date >= start,
|
||||||
StockDailyModel.trade_date <= end,
|
StockDailyModel.trade_date <= end,
|
||||||
|
StockDailyModel.adjust == adjust, # 研究主口径:不复权(v2 §8)
|
||||||
)
|
)
|
||||||
.order_by(StockDailyModel.trade_date)
|
.order_by(StockDailyModel.trade_date)
|
||||||
).all()
|
)
|
||||||
|
rows = self._session.scalars(stmt).all()
|
||||||
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
||||||
|
|
||||||
def stream_range_many_columns(
|
def stream_range_many_columns(
|
||||||
@@ -176,6 +184,7 @@ class SqlAlchemyDailyBarRepository:
|
|||||||
start: date,
|
start: date,
|
||||||
end: date,
|
end: date,
|
||||||
columns: Sequence[str],
|
columns: Sequence[str],
|
||||||
|
adjust: str = "none",
|
||||||
) -> Iterator[tuple]:
|
) -> Iterator[tuple]:
|
||||||
"""流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。
|
"""流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。
|
||||||
|
|
||||||
@@ -193,6 +202,7 @@ class SqlAlchemyDailyBarRepository:
|
|||||||
StockDailyModel.symbol.in_(list(symbols)),
|
StockDailyModel.symbol.in_(list(symbols)),
|
||||||
StockDailyModel.trade_date >= start,
|
StockDailyModel.trade_date >= start,
|
||||||
StockDailyModel.trade_date <= end,
|
StockDailyModel.trade_date <= end,
|
||||||
|
StockDailyModel.adjust == adjust,
|
||||||
)
|
)
|
||||||
.order_by(StockDailyModel.symbol, StockDailyModel.trade_date)
|
.order_by(StockDailyModel.symbol, StockDailyModel.trade_date)
|
||||||
.execution_options(yield_per=20000)
|
.execution_options(yield_per=20000)
|
||||||
|
|||||||
@@ -75,6 +75,7 @@ def load_daily_df(
|
|||||||
start: date,
|
start: date,
|
||||||
end: date,
|
end: date,
|
||||||
columns: list[str],
|
columns: list[str],
|
||||||
|
adjust: str = "none",
|
||||||
) -> pd.DataFrame:
|
) -> pd.DataFrame:
|
||||||
"""从 Repository 装配行情长表(供研究/选股共用)。
|
"""从 Repository 装配行情长表(供研究/选股共用)。
|
||||||
|
|
||||||
@@ -86,12 +87,14 @@ def load_daily_df(
|
|||||||
streamer = getattr(daily_repo, "stream_range_many_columns", None)
|
streamer = getattr(daily_repo, "stream_range_many_columns", None)
|
||||||
if streamer is not None:
|
if streamer is not None:
|
||||||
try:
|
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 实现)
|
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
|
||||||
pass
|
pass
|
||||||
get_many = getattr(daily_repo, "get_range_many", None)
|
get_many = getattr(daily_repo, "get_range_many", None)
|
||||||
if get_many is not None:
|
if get_many is not None:
|
||||||
bars = list(get_many(symbols, start, end))
|
bars = list(get_many(symbols, start, end, adjust=adjust))
|
||||||
else: # 兜底:逐只查询
|
else: # 兜底:逐只查询
|
||||||
bars = []
|
bars = []
|
||||||
for sym in symbols:
|
for sym in symbols:
|
||||||
@@ -134,5 +137,10 @@ class ResearchService:
|
|||||||
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
|
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
|
||||||
required = self._engine.required_columns(spec)
|
required = self._engine.required_columns(spec)
|
||||||
return load_daily_df(
|
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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ def client(tmp_path) -> TestClient:
|
|||||||
bars = bars_dataframe_to_daily_bars(daily_df)
|
bars = bars_dataframe_to_daily_bars(daily_df)
|
||||||
|
|
||||||
class _MemDailyRepo:
|
class _MemDailyRepo:
|
||||||
def get_range_many(self, symbols, start, end):
|
def get_range_many(self, symbols, start, end, adjust="none"):
|
||||||
out = []
|
out = []
|
||||||
for b in bars:
|
for b in bars:
|
||||||
if b.symbol in symbols and start <= b.trade_date <= end:
|
if b.symbol in symbols and start <= b.trade_date <= end:
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ class _FakeDailyRepo:
|
|||||||
def __init__(self, bars):
|
def __init__(self, bars):
|
||||||
self._bars = bars
|
self._bars = bars
|
||||||
|
|
||||||
def get_range_many(self, symbols, start, end):
|
def get_range_many(self, symbols, start, end, adjust="none"):
|
||||||
return [b for b in self._bars if b.symbol in symbols and start <= b.trade_date <= end]
|
return [b for b in self._bars if b.symbol in symbols and start <= b.trade_date <= end]
|
||||||
|
|
||||||
def get_range(self, symbol, start, end):
|
def get_range(self, symbol, start, end):
|
||||||
|
|||||||
@@ -0,0 +1,106 @@
|
|||||||
|
"""M7.3 行情口径测试:Repository 读路径按 adjust 过滤(不复权主口径),spec 记录口径。
|
||||||
|
|
||||||
|
混合行场景:同 symbol/date 存在 tushare/none 与 sina/qfq 行时,研究读取
|
||||||
|
(get_range_many / stream)默认只取 adjust=none —— 消除「混合口径污染因子」风险。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from app.domain.entities.market import DailyBar
|
||||||
|
from app.domain.entities.research import ResearchSpec
|
||||||
|
from app.domain.entities.selection import SelectionQuery
|
||||||
|
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||||
|
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||||
|
SqlAlchemyDailyBarRepository,
|
||||||
|
)
|
||||||
|
from sqlalchemy import create_engine
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
_D = date(2024, 6, 3)
|
||||||
|
_D2 = date(2024, 6, 4)
|
||||||
|
_D3 = date(2024, 6, 5)
|
||||||
|
|
||||||
|
|
||||||
|
def _bar(adjust: str, close: str, source: str = "tushare", day=None) -> DailyBar:
|
||||||
|
return DailyBar(
|
||||||
|
symbol="600519.SH",
|
||||||
|
trade_date=day or _D,
|
||||||
|
source=source,
|
||||||
|
adjust=adjust,
|
||||||
|
open=Decimal("100"), high=Decimal("101"), low=Decimal("99"),
|
||||||
|
close=Decimal(close), volume=Decimal("1000"), amount=Decimal("100000"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestAdjustFilter:
|
||||||
|
def _session(self, tmp_path) -> Session:
|
||||||
|
engine = create_engine(f"sqlite:///{tmp_path / 'adj.db'}", future=True)
|
||||||
|
Base.metadata.create_all(engine)
|
||||||
|
return Session(engine)
|
||||||
|
|
||||||
|
def test_get_range_many_filters_adjust(self, tmp_path) -> None:
|
||||||
|
with self._session(tmp_path) as session:
|
||||||
|
repo = SqlAlchemyDailyBarRepository(session)
|
||||||
|
# 唯一键 (symbol, trade_date):同键共存不可能 —— 用连续三天模拟
|
||||||
|
# none 主口径两天 + sina/qfq 兜底一天
|
||||||
|
repo.upsert_many(
|
||||||
|
[
|
||||||
|
_bar("none", "1700", day=_D),
|
||||||
|
_bar("none", "1710", day=_D2),
|
||||||
|
_bar("qfq", "1680", source="sina", day=_D3),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
session.commit()
|
||||||
|
|
||||||
|
none_rows = repo.get_range_many(["600519.SH"], _D, _D3, adjust="none")
|
||||||
|
assert len(none_rows) == 2 and {float(r.close) for r in none_rows} == {1700, 1710}
|
||||||
|
qfq_rows = repo.get_range_many(["600519.SH"], _D, _D3, adjust="qfq")
|
||||||
|
assert len(qfq_rows) == 1 and float(qfq_rows[0].close) == 1680
|
||||||
|
|
||||||
|
def test_stream_filters_adjust(self, tmp_path) -> None:
|
||||||
|
with self._session(tmp_path) as session:
|
||||||
|
repo = SqlAlchemyDailyBarRepository(session)
|
||||||
|
repo.upsert_many(
|
||||||
|
[
|
||||||
|
_bar("none", "1700", day=_D),
|
||||||
|
_bar("none", "1710", day=_D2),
|
||||||
|
_bar("qfq", "1680", source="sina", day=_D3),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
session.commit()
|
||||||
|
rows = list(
|
||||||
|
repo.stream_range_many_columns(
|
||||||
|
["600519.SH"], _D, _D3, ["close"], adjust="none"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert len(rows) == 2 and {float(r[-1]) for r in rows} == {1700, 1710}
|
||||||
|
qrows = list(
|
||||||
|
repo.stream_range_many_columns(
|
||||||
|
["600519.SH"], _D, _D3, ["close"], adjust="qfq"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert len(qrows) == 1 and float(qrows[0][-1]) == 1680
|
||||||
|
|
||||||
|
|
||||||
|
class TestSpecRecordsAdjustment:
|
||||||
|
def test_research_spec_default_and_field(self) -> None:
|
||||||
|
spec = ResearchSpec(
|
||||||
|
type="backtest",
|
||||||
|
factors=[{"name": "momentum_60", "weight": 1}],
|
||||||
|
period=(date(2024, 1, 1), date(2024, 6, 1)),
|
||||||
|
)
|
||||||
|
assert spec.price_adjustment == "none"
|
||||||
|
snap = spec.model_dump(mode="json")
|
||||||
|
assert snap["price_adjustment"] == "none" # 结果 config_snapshot 可溯源
|
||||||
|
|
||||||
|
def test_selection_query_adjustment_in_snapshot(self) -> None:
|
||||||
|
q = SelectionQuery(factors=[{"name": "momentum_60", "weight": 1}], top_n=5)
|
||||||
|
assert q.price_adjustment == "none"
|
||||||
|
assert q.model_dump()["price_adjustment"] == "none"
|
||||||
|
q2 = SelectionQuery(
|
||||||
|
factors=[{"name": "momentum_60", "weight": 1}], top_n=5, price_adjustment="qfq"
|
||||||
|
)
|
||||||
|
assert q2.price_adjustment == "qfq"
|
||||||
@@ -69,7 +69,7 @@ class _MemDailyRepo:
|
|||||||
def get_range(self, symbol, start, end) -> list[DailyBar]:
|
def get_range(self, symbol, start, end) -> list[DailyBar]:
|
||||||
return self._bars([symbol], start, end)
|
return self._bars([symbol], start, end)
|
||||||
|
|
||||||
def get_range_many(self, symbols, start, end) -> list[DailyBar]:
|
def get_range_many(self, symbols, start, end, adjust="none") -> list[DailyBar]:
|
||||||
return self._bars(list(symbols), start, end)
|
return self._bars(list(symbols), start, end)
|
||||||
|
|
||||||
def latest_date(self, symbol: str) -> date | None:
|
def latest_date(self, symbol: str) -> date | None:
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ class _MemDailyRepo:
|
|||||||
def get_range(self, symbol, start, end):
|
def get_range(self, symbol, start, end):
|
||||||
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
|
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
|
||||||
|
|
||||||
def get_range_many(self, symbols, start, end):
|
def get_range_many(self, symbols, start, end, adjust="none"):
|
||||||
syms = set(symbols)
|
syms = set(symbols)
|
||||||
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
|
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
|
||||||
|
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ class _MemDailyRepo:
|
|||||||
def get_range(self, symbol, start, end):
|
def get_range(self, symbol, start, end):
|
||||||
return [b for b in self._bars_all if b.symbol == symbol and start <= b.trade_date <= end]
|
return [b for b in self._bars_all if b.symbol == symbol and start <= b.trade_date <= end]
|
||||||
|
|
||||||
def get_range_many(self, symbols, start, end):
|
def get_range_many(self, symbols, start, end, adjust="none"):
|
||||||
syms = set(symbols)
|
syms = set(symbols)
|
||||||
return [b for b in self._bars_all if b.symbol in syms and start <= b.trade_date <= end]
|
return [b for b in self._bars_all if b.symbol in syms and start <= b.trade_date <= end]
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user