feat(portfolio): P2/P3 策略层向量化 + fundamentals批量提速解锁长回测
P2 行情向量化(get_price→get_closes_panel,口径实证 max_abs_diff=0.0 零偏差):
- momentum_timing: _cal_rps/_select_stocks/_cal_buy_sign 三处向量化
- small_cap: _cal_momentum_score 用 close.min/max 代理 low/high(方案A)
P3 fundamentals 批量(small_cap _pick_stocks 加 fields=[market_cap,eps],
对接数据session f416a17 get_fundamentals_df 按需短路):
- 02 000985 全市场 5128只 ~19min卡死 → 195s 跑通解锁
验收: 72单测全过; VPS 03短回测+138%(口径与get_price一致diff=0)/
02全市场-11%(2024Q1小盘股灾期合理)/01沪深300可跑
This commit is contained in:
@@ -210,6 +210,9 @@ class MomentumTimingStrategy:
|
||||
改为取 ``preDate ~ curDate`` 区间算**百分比涨跌幅**(更符合 RPS 语义,
|
||||
原代码用绝对差值排序会偏向高价股,见 notes.md「移植记录」)。
|
||||
|
||||
性能改造:用 ``get_closes_panel`` 一次批量取宽表(替代 get_price+pivot,
|
||||
5128 只从逐只循环秒级降到 UNION ALL 批量)。数值口径不变(fq='raw')。
|
||||
|
||||
Returns:
|
||||
DataFrame[code, rps_value],按 rps_value 降序;``rps_value = 99 - 100*i/n``。
|
||||
"""
|
||||
@@ -217,31 +220,18 @@ class MomentumTimingStrategy:
|
||||
if n == 0:
|
||||
return pd.DataFrame({"code": [], "rps_value": []})
|
||||
try:
|
||||
df = self.provider.get_price(
|
||||
stocks,
|
||||
start_date=pre_date,
|
||||
end_date=cur_date,
|
||||
frequency="daily",
|
||||
fields=["close"],
|
||||
panel=False,
|
||||
fill_paused=False,
|
||||
panel = self.provider.get_closes_panel(
|
||||
stocks, pre_date, cur_date, fq="raw",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("_cal_rps get_price 失败: %s", exc)
|
||||
logger.warning("_cal_rps get_closes_panel 失败: %s", exc)
|
||||
return pd.DataFrame({"code": [], "rps_value": []})
|
||||
if df is None or df.empty:
|
||||
return pd.DataFrame({"code": [], "rps_value": []})
|
||||
try:
|
||||
pivot = df.pivot(index="time", columns="code", values="close")
|
||||
except Exception as exc:
|
||||
logger.warning("_cal_rps pivot 失败: %s", exc)
|
||||
return pd.DataFrame({"code": [], "rps_value": []})
|
||||
if pivot.empty or len(pivot) < 2:
|
||||
if panel is None or panel.empty or len(panel) < 2:
|
||||
return pd.DataFrame({"code": [], "rps_value": []})
|
||||
|
||||
# 每只股票涨跌幅(末值/首值 - 1)
|
||||
first = pivot.iloc[0]
|
||||
last = pivot.iloc[-1]
|
||||
# 每只股票涨跌幅(末值/首值 - 1) — 向量化
|
||||
first = panel.iloc[0]
|
||||
last = panel.iloc[-1]
|
||||
with np.errstate(divide="ignore", invalid="ignore"):
|
||||
returns = (last / first) - 1.0
|
||||
# 过滤 NaN/Inf(数据不全或首值为 0)
|
||||
@@ -282,46 +272,43 @@ class MomentumTimingStrategy:
|
||||
"""均线动量过滤:``close > MA_short`` 且 ``MA_short > MA_long``。
|
||||
|
||||
原策略 ``data[security].mavg(5,'close')`` (聚宽 Security.mavg),
|
||||
翻译为 provider.get_price(count=ma_long) 后段求均值。
|
||||
翻译为 ``get_closes_panel`` 批量取 close 宽表,向量化算双均线。
|
||||
|
||||
性能改造:用 ``get_closes_panel`` 替代 ``get_price(count=N)``。
|
||||
count=N → ``start = cur - N*2 自然日``、``end = cur``、宽表 ``.tail(N)`` 切片
|
||||
(避免自然日 vs 交易日的换算歧义,``.tail(N)`` 最稳)。数值口径不变(fq='raw')。
|
||||
"""
|
||||
cfg = self.config
|
||||
if not stocks:
|
||||
return []
|
||||
start_date = _shift_date(cur_date, -cfg.ma_long * 2)
|
||||
try:
|
||||
df = self.provider.get_price(
|
||||
stocks,
|
||||
end_date=cur_date,
|
||||
frequency="daily",
|
||||
fields=["close"],
|
||||
count=cfg.ma_long,
|
||||
panel=False,
|
||||
fill_paused=False,
|
||||
panel = self.provider.get_closes_panel(
|
||||
stocks, start_date, cur_date, fq="raw",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("_select_stocks get_price 失败: %s", exc)
|
||||
logger.warning("_select_stocks get_closes_panel 失败: %s", exc)
|
||||
return []
|
||||
if df is None or df.empty:
|
||||
if panel is None or panel.empty:
|
||||
return []
|
||||
try:
|
||||
pivot = df.pivot(index="time", columns="code", values="close")
|
||||
except Exception:
|
||||
return []
|
||||
if pivot.empty:
|
||||
panel = panel.tail(cfg.ma_long)
|
||||
if panel.empty:
|
||||
return []
|
||||
|
||||
out: List[str] = []
|
||||
for col in pivot.columns:
|
||||
series = pivot[col].dropna()
|
||||
if len(series) < cfg.ma_long:
|
||||
continue
|
||||
close = float(series.iloc[-1])
|
||||
ma_short = float(series.tail(cfg.ma_short).mean())
|
||||
ma_long = float(series.tail(cfg.ma_long).mean())
|
||||
if np.isnan(close) or np.isnan(ma_short) or np.isnan(ma_long):
|
||||
continue
|
||||
if close > ma_short and ma_short > ma_long:
|
||||
out.append(col)
|
||||
return out
|
||||
# 向量化:对每列算 close/ma_short/ma_long,过滤条件 close>ma_short>ma_long
|
||||
valid_count = panel.notna().sum() # 每列非 NaN 计数,等价于原 dropna 后长度
|
||||
close = panel.iloc[-1]
|
||||
ma_short = panel.tail(cfg.ma_short).mean()
|
||||
ma_long = panel.mean() # panel 已 tail(ma_long),整体均值即 MA_long
|
||||
mask = (
|
||||
(valid_count >= cfg.ma_long)
|
||||
& close.notna()
|
||||
& ma_short.notna()
|
||||
& ma_long.notna()
|
||||
& (close > ma_short)
|
||||
& (ma_short > ma_long)
|
||||
)
|
||||
return list(mask.index[mask])
|
||||
|
||||
# =================== calBuySign (牛熊分界) ===================
|
||||
def _cal_buy_sign(
|
||||
@@ -333,8 +320,12 @@ class MomentumTimingStrategy:
|
||||
"""统计 past_day 均线上方的指数占比 > index_thre → 牛市(True)。
|
||||
|
||||
原策略 'index' 模式(第 110-115 行):对每个指数算 ``mavg(past_day,'close')``
|
||||
与 ``mavg(1,'close')`` 比较。翻译为取 past_day 日 close(含当日),
|
||||
算均值与最后一根 close 比较。
|
||||
与 ``mavg(1,'close')`` 比较。翻译为 ``get_closes_panel`` 批量取 close 宽表,
|
||||
向量化算 past_day 均线。
|
||||
|
||||
性能改造:用 ``get_closes_panel`` 替代 ``get_price(count=N)``。
|
||||
count=N → ``start = cur - N*2 自然日``、``end = cur``、宽表 ``.tail(N)`` 切片。
|
||||
数值口径不变(fq='raw')。
|
||||
|
||||
⚠️ 原代码 ``float(count)/len(indexList)`` 在 py2 是浮点除法(因 float()强转),
|
||||
与 py3 一致。这里保留浮点除法语义。
|
||||
@@ -342,39 +333,31 @@ class MomentumTimingStrategy:
|
||||
cfg = self.config
|
||||
if not index_list:
|
||||
return False
|
||||
start_date = _shift_date(cur_date, -past_day * 2)
|
||||
try:
|
||||
df = self.provider.get_price(
|
||||
index_list,
|
||||
end_date=cur_date,
|
||||
frequency="daily",
|
||||
fields=["close"],
|
||||
count=past_day,
|
||||
panel=False,
|
||||
fill_paused=False,
|
||||
panel = self.provider.get_closes_panel(
|
||||
index_list, start_date, cur_date, fq="raw",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("_cal_buy_sign get_price 失败: %s", exc)
|
||||
logger.warning("_cal_buy_sign get_closes_panel 失败: %s", exc)
|
||||
return False
|
||||
if df is None or df.empty:
|
||||
if panel is None or panel.empty:
|
||||
return False
|
||||
try:
|
||||
pivot = df.pivot(index="time", columns="code", values="close")
|
||||
except Exception:
|
||||
return False
|
||||
if pivot.empty:
|
||||
panel = panel.tail(past_day)
|
||||
if panel.empty:
|
||||
return False
|
||||
|
||||
count = 0
|
||||
for col in pivot.columns:
|
||||
series = pivot[col].dropna()
|
||||
if len(series) < 2:
|
||||
continue
|
||||
ma_past = float(series.tail(past_day).mean())
|
||||
cur_close = float(series.iloc[-1])
|
||||
if np.isnan(ma_past) or np.isnan(cur_close):
|
||||
continue
|
||||
if cur_close > ma_past:
|
||||
count += 1
|
||||
# 向量化:对每列算 cur_close 与 past_day 均值,统计 close > ma_past 的占比
|
||||
valid_count = panel.notna().sum()
|
||||
cur_close = panel.iloc[-1]
|
||||
ma_past = panel.mean()
|
||||
mask = (
|
||||
(valid_count >= 2)
|
||||
& cur_close.notna()
|
||||
& ma_past.notna()
|
||||
& (cur_close > ma_past)
|
||||
)
|
||||
count = int(mask.sum())
|
||||
return (count / len(index_list)) > cfg.index_thre
|
||||
|
||||
# =================== 调仓辅助 ===================
|
||||
@@ -418,4 +401,17 @@ def _to_date_str(value: Any) -> str:
|
||||
return str(value)[:10]
|
||||
|
||||
|
||||
def _shift_date(date_str: str, days: int) -> str:
|
||||
"""字符串日期加减天数,返回 YYYY-MM-DD。
|
||||
|
||||
用于 ``count=N`` → ``start = end - N*2 自然日`` 的换算(配合宽表 ``.tail(N)`` 切片,
|
||||
避免自然日 vs 交易日的歧义)。
|
||||
"""
|
||||
try:
|
||||
dt = datetime.datetime.strptime(date_str[:10], "%Y-%m-%d")
|
||||
except (ValueError, TypeError):
|
||||
return date_str
|
||||
return (dt + datetime.timedelta(days=days)).strftime("%Y-%m-%d")
|
||||
|
||||
|
||||
__all__ = ["MomentumTimingStrategy", "MomentumTimingConfig"]
|
||||
|
||||
@@ -28,6 +28,7 @@
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, List, Optional
|
||||
@@ -191,8 +192,14 @@ class SmallCapStrategy:
|
||||
return []
|
||||
|
||||
# 2) get_fundamentals_df 一次性取 market_cap + eps
|
||||
# fields= 按需短路源表(P3 批量提速,数据session commit f416a17:
|
||||
# 5128 只 ~19min→~2min)。_pick_stocks 只用 market_cap(排序)+eps(>0过滤)。
|
||||
# ⚠️ 若以后给 _pick_stocks 加新过滤(ROE/营收等),必须把列名加进 fields=,
|
||||
# 否则该列返 NaN→过滤静默失效;fields=None 仍全列(向后兼容但慢)。
|
||||
try:
|
||||
df = self.provider.get_fundamentals_df(candidates, date=previous_date)
|
||||
df = self.provider.get_fundamentals_df(
|
||||
candidates, date=previous_date, fields=["market_cap", "eps"],
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("get_fundamentals_df 失败: %s", exc)
|
||||
return []
|
||||
@@ -260,6 +267,12 @@ class SmallCapStrategy:
|
||||
- ``score = (cur-low_130) + (cur-high_130) + (cur-avg_15)``
|
||||
- 升序(分数越低越靠前:price 接近 130 日低 / 低于均线 → 偏底部)
|
||||
|
||||
性能改造(决策方案 A):用 ``get_closes_panel`` 批量取 close 宽表向量化,
|
||||
**用 close.rolling(130).min/max 代理 low.min()/high.max()**。
|
||||
原因:``get_closes_panel`` 只返 close(不含 high/low);用 close 极值代理是有意
|
||||
决策(spec 明确允许)——对动量评分的"底部反弹偏好"语义无实质影响(都衡量
|
||||
当前价在 130 日极值区间的位置),换 5128 只逐只循环 → 一次批量(33s → 秒级)。
|
||||
|
||||
py2→py3:``df.sort(columns=)`` → ``df.sort_values(by=)``。
|
||||
|
||||
Returns:
|
||||
@@ -269,49 +282,36 @@ class SmallCapStrategy:
|
||||
if not stocks:
|
||||
return pd.DataFrame(columns=["score"])
|
||||
|
||||
# 一次性取 ma_window=130 日 close/high/low(对所有候选)
|
||||
# 一次性取 ma_window=130 日 close(批量宽表,用 min/max 代理 low/high)
|
||||
start_date = _shift_date(end_date, -cfg.ma_window * 2)
|
||||
try:
|
||||
df = self.provider.get_price(
|
||||
stocks,
|
||||
end_date=end_date,
|
||||
frequency="daily",
|
||||
fields=["close", "high", "low"],
|
||||
count=cfg.ma_window,
|
||||
panel=False,
|
||||
fill_paused=False,
|
||||
panel = self.provider.get_closes_panel(
|
||||
stocks, start_date, end_date, fq="raw",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("_cal_momentum_score get_price 失败: %s", exc)
|
||||
logger.warning("_cal_momentum_score get_closes_panel 失败: %s", exc)
|
||||
return pd.DataFrame(columns=["score"])
|
||||
if df is None or df.empty:
|
||||
if panel is None or panel.empty:
|
||||
return pd.DataFrame(columns=["score"])
|
||||
panel = panel.tail(cfg.ma_window)
|
||||
if panel.empty:
|
||||
return pd.DataFrame(columns=["score"])
|
||||
|
||||
scores: dict[str, float] = {}
|
||||
for code in stocks:
|
||||
sub = df[df["code"] == code] if "code" in df.columns else df
|
||||
if sub is None or sub.empty:
|
||||
continue
|
||||
close_series = sub["close"].dropna() if "close" in sub.columns else None
|
||||
high_series = sub["high"].dropna() if "high" in sub.columns else None
|
||||
low_series = sub["low"].dropna() if "low" in sub.columns else None
|
||||
if close_series is None or close_series.empty:
|
||||
continue
|
||||
cur_price = float(close_series.iloc[-1])
|
||||
if not np.isfinite(cur_price):
|
||||
continue
|
||||
# 130 日最低 / 最高(skip_paused=True 后 dropna)
|
||||
low_130 = float(low_series.min()) if low_series is not None and not low_series.empty else cur_price
|
||||
high_130 = float(high_series.max()) if high_series is not None and not high_series.empty else cur_price
|
||||
# 15 日均线:close 序列最后 15 根均值
|
||||
ma15 = float(close_series.tail(cfg.ma_short).mean()) if len(close_series) >= 1 else cur_price
|
||||
if not (np.isfinite(low_130) and np.isfinite(high_130) and np.isfinite(ma15)):
|
||||
continue
|
||||
score = (cur_price - low_130) + (cur_price - high_130) + (cur_price - ma15)
|
||||
scores[code] = score
|
||||
# 向量化算 score = (cur-low) + (cur-high) + (cur-ma15)
|
||||
# 低/高用 close 序列代理(原 high.max()/low.min())
|
||||
cur_price = panel.iloc[-1]
|
||||
low_proxy = panel.min() # 130 日 close 最低(代理 low.min())
|
||||
high_proxy = panel.max() # 130 日 close 最高(代理 high.max())
|
||||
ma15 = panel.tail(cfg.ma_short).mean()
|
||||
|
||||
if not scores:
|
||||
# 过滤:cur_price 必须有效 + 至少 1 个有效值(原代码 close_series.empty 跳过)
|
||||
valid_count = panel.notna().sum()
|
||||
score = (cur_price - low_proxy) + (cur_price - high_proxy) + (cur_price - ma15)
|
||||
mask = (valid_count >= 1) & cur_price.notna() & np.isfinite(cur_price)
|
||||
score = score[mask].dropna()
|
||||
if score.empty:
|
||||
return pd.DataFrame(columns=["score"])
|
||||
out = pd.DataFrame.from_dict(scores, orient="index", columns=["score"])
|
||||
out = score.to_frame("score")
|
||||
# 升序:分数越低越靠前(原策略 df.sort(columns='score', ascending=True))
|
||||
out = out.sort_values("score", ascending=True)
|
||||
return out
|
||||
@@ -407,4 +407,17 @@ def _is_valid_positive_number(v: Any) -> bool:
|
||||
return fv > 0
|
||||
|
||||
|
||||
def _shift_date(date_str: str, days: int) -> str:
|
||||
"""字符串日期加减天数,返回 YYYY-MM-DD。
|
||||
|
||||
用于 ``count=N`` → ``start = end - N*2 自然日`` 的换算(配合宽表 ``.tail(N)`` 切片,
|
||||
避免自然日 vs 交易日的歧义)。
|
||||
"""
|
||||
try:
|
||||
dt = datetime.datetime.strptime(date_str[:10], "%Y-%m-%d")
|
||||
except (ValueError, TypeError):
|
||||
return date_str
|
||||
return (dt + datetime.timedelta(days=days)).strftime("%Y-%m-%d")
|
||||
|
||||
|
||||
__all__ = ["SmallCapStrategy", "SmallCapConfig"]
|
||||
|
||||
@@ -2,6 +2,9 @@
|
||||
|
||||
策略层只测**逻辑分支正确**(RPS / 均线 / 牛熊信号 / 调仓),不测真实数据。
|
||||
真实数据回测在 VPS 跑,这里只保证策略翻译等价 + 两个原始 bug 已修复。
|
||||
|
||||
性能改造后:策略用 ``get_closes_panel`` 批量取 close 宽表(替代 get_price+pivot),
|
||||
测试 mock 同步切到 ``get_closes_panel`` 返回宽表 DataFrame。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -25,13 +28,13 @@ from tests.portfolio.conftest import FakeContext, FakePosition
|
||||
def make_strategy(
|
||||
*,
|
||||
index_stocks_map: Optional[Dict[str, List[str]]] = None,
|
||||
price_df_map: Optional[Dict[Any, pd.DataFrame]] = None,
|
||||
panel_map: Optional[Dict[Any, pd.DataFrame]] = None,
|
||||
config: Optional[MomentumTimingConfig] = None,
|
||||
) -> MomentumTimingStrategy:
|
||||
"""构造一个 mock provider + mock broker 驱动的策略。
|
||||
|
||||
- index_stocks_map: get_index_stocks 返回,dict[index] -> List[code]
|
||||
- price_df_map: get_price 按 (security, fields, count) 或 (security, start, end) 缓存的返回
|
||||
- panel_map: get_closes_panel 按 (symbols_tuple) 或 (symbols_tuple, start, end) 缓存的宽表返回
|
||||
"""
|
||||
provider = MagicMock(name="provider")
|
||||
|
||||
@@ -59,29 +62,28 @@ def make_strategy(
|
||||
"high_limit": 11.0, "low_limit": 9.0,
|
||||
}
|
||||
|
||||
# get_price 按 key 缓存(支持 count 模式 + start/end 模式)
|
||||
# 规范化:把 key 第一项(list)转 tuple 以保证可 hash
|
||||
def _normalize_key(k: Any) -> Any:
|
||||
# get_closes_panel 按 key 缓存:支持精确 (tuple(symbols), start, end) 与
|
||||
# 通配 (tuple(symbols),) 两种 key(精确优先,fallback 忽略日期)
|
||||
def _normalize_syms(k: Any) -> Any:
|
||||
if isinstance(k, tuple) and k and isinstance(k[0], (list, tuple)):
|
||||
return (tuple(k[0]),) + tuple(k[1:])
|
||||
return k
|
||||
|
||||
price_df_map = {_normalize_key(k): v for k, v in (price_df_map or {}).items()}
|
||||
panel_map = {_normalize_syms(k): v for k, v in (panel_map or {}).items()}
|
||||
|
||||
def _get_price(security, **kwargs):
|
||||
# 构造 cache key:两种取数模式
|
||||
# 1) count 模式:(sec_key, fields, count)
|
||||
# 2) start/end 模式:(sec_key, fields, start_date, end_date)
|
||||
# 注意:list 不可 hash → 转 tuple
|
||||
sec_key = tuple(security) if isinstance(security, list) else security
|
||||
fields = tuple(kwargs.get("fields") or [])
|
||||
if kwargs.get("count") is not None:
|
||||
key = (sec_key, fields, kwargs.get("count"))
|
||||
else:
|
||||
key = (sec_key, fields, kwargs.get("start_date"), kwargs.get("end_date"))
|
||||
return price_df_map.get(key, pd.DataFrame())
|
||||
def _get_closes_panel(symbols, start=None, end=None, interval="d", fq="raw"):
|
||||
syms_key = tuple(symbols) if isinstance(symbols, list) else symbols
|
||||
# 精确 key 优先
|
||||
exact = panel_map.get((syms_key, start, end))
|
||||
if exact is not None:
|
||||
return exact
|
||||
# fallback: 忽略 start/end
|
||||
fallback = panel_map.get((syms_key,))
|
||||
if fallback is not None:
|
||||
return fallback
|
||||
return pd.DataFrame(index=pd.DatetimeIndex([]))
|
||||
|
||||
provider.get_price.side_effect = _get_price
|
||||
provider.get_closes_panel.side_effect = _get_closes_panel
|
||||
|
||||
broker = BrokerFacade()
|
||||
broker.order_target_value = MagicMock(return_value=MagicMock(filled=100))
|
||||
@@ -94,29 +96,30 @@ def make_strategy(
|
||||
return MomentumTimingStrategy(provider=provider, broker=broker, config=config)
|
||||
|
||||
|
||||
def _make_close_panel(
|
||||
def _make_close_wide(
|
||||
codes: List[str],
|
||||
closes: List[List[float]],
|
||||
end_date: str = "2024-09-30",
|
||||
days: int = 30,
|
||||
) -> pd.DataFrame:
|
||||
"""构造 panel=False 风格的 close DataFrame。
|
||||
"""构造 get_closes_panel 风格的 close 宽表 (index=DatetimeIndex, columns=codes)。
|
||||
|
||||
Args:
|
||||
codes: 股票代码列表
|
||||
codes: 股票代码列表(列名)
|
||||
closes: 每只股票的 close 序列(长度 <= days, 不足重复首值)
|
||||
end_date: 最后一根 K 线日期
|
||||
days: 总 K 线根数(默认 30)
|
||||
"""
|
||||
end_dt = datetime.strptime(end_date, "%Y-%m-%d")
|
||||
dates = [(end_dt - timedelta(days=days - 1 - i)).strftime("%Y-%m-%d") for i in range(days)]
|
||||
rows = []
|
||||
dates = pd.DatetimeIndex([
|
||||
end_dt - timedelta(days=days - 1 - i) for i in range(days)
|
||||
])
|
||||
data: Dict[str, List[float]] = {}
|
||||
for code, close_list in zip(codes, closes):
|
||||
# 不足 days 的补首值
|
||||
full = list(close_list) + [close_list[-1]] * (days - len(close_list))
|
||||
for d, c in zip(dates, full):
|
||||
rows.append({"time": pd.Timestamp(d), "code": code, "close": float(c)})
|
||||
return pd.DataFrame(rows)
|
||||
data[code] = [float(c) for c in full]
|
||||
return pd.DataFrame(data, index=dates)
|
||||
|
||||
|
||||
# =================== initialize ===================
|
||||
@@ -153,17 +156,16 @@ class TestCalRps:
|
||||
# 3 只股票,涨幅依次为 +100% / +50% / 0%
|
||||
# preDate 首值 = 10, curDate 末值 = 20 / 15 / 10
|
||||
codes = ["A.XSHG", "B.XSHG", "C.XSHG"]
|
||||
df = pd.DataFrame([
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "A.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "A.XSHG", "close": 20.0}, # +100%
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "B.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "B.XSHG", "close": 15.0}, # +50%
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "C.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "C.XSHG", "close": 10.0}, # 0%
|
||||
])
|
||||
s = make_strategy(price_df_map={
|
||||
# 按 start_date/end_date 取数,确认 _cal_rps 走的是区间查询
|
||||
((str(codes),) if False else (tuple(codes), ("close",), "2024-09-01", "2024-09-30")): df,
|
||||
panel = pd.DataFrame(
|
||||
{
|
||||
"A.XSHG": [10.0, 20.0], # +100%
|
||||
"B.XSHG": [10.0, 15.0], # +50%
|
||||
"C.XSHG": [10.0, 10.0], # 0%
|
||||
},
|
||||
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
|
||||
)
|
||||
s = make_strategy(panel_map={
|
||||
(tuple(codes), "2024-09-01", "2024-09-30"): panel,
|
||||
})
|
||||
out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01")
|
||||
|
||||
@@ -177,14 +179,15 @@ class TestCalRps:
|
||||
def test_rps_descending_by_return(self):
|
||||
"""涨幅大的排前(降序)。"""
|
||||
codes = ["X.XSHG", "Y.XSHG"]
|
||||
df = pd.DataFrame([
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "X.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "X.XSHG", "close": 12.0}, # +20%
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "Y.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "Y.XSHG", "close": 15.0}, # +50%
|
||||
])
|
||||
s = make_strategy(price_df_map={
|
||||
(tuple(codes), ("close",), "2024-09-01", "2024-09-30"): df,
|
||||
panel = pd.DataFrame(
|
||||
{
|
||||
"X.XSHG": [10.0, 12.0], # +20%
|
||||
"Y.XSHG": [10.0, 15.0], # +50%
|
||||
},
|
||||
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
|
||||
)
|
||||
s = make_strategy(panel_map={
|
||||
(tuple(codes), "2024-09-01", "2024-09-30"): panel,
|
||||
})
|
||||
out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01")
|
||||
# Y 涨幅大,排前
|
||||
@@ -194,16 +197,16 @@ class TestCalRps:
|
||||
def test_rps_filters_nan_and_zero_first(self):
|
||||
"""首值为 0(除零)或 NaN → 过滤掉。"""
|
||||
codes = ["GOOD.XSHG", "ZERO.XSHG", "NAN.XSHG"]
|
||||
df = pd.DataFrame([
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "GOOD.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "GOOD.XSHG", "close": 20.0},
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "ZERO.XSHG", "close": 0.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "ZERO.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "NAN.XSHG", "close": np.nan},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "NAN.XSHG", "close": 10.0},
|
||||
])
|
||||
s = make_strategy(price_df_map={
|
||||
(tuple(codes), ("close",), "2024-09-01", "2024-09-30"): df,
|
||||
panel = pd.DataFrame(
|
||||
{
|
||||
"GOOD.XSHG": [10.0, 20.0],
|
||||
"ZERO.XSHG": [0.0, 10.0],
|
||||
"NAN.XSHG": [np.nan, 10.0],
|
||||
},
|
||||
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
|
||||
)
|
||||
s = make_strategy(panel_map={
|
||||
(tuple(codes), "2024-09-01", "2024-09-30"): panel,
|
||||
})
|
||||
out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01")
|
||||
assert list(out["code"]) == ["GOOD.XSHG"]
|
||||
@@ -219,10 +222,9 @@ class TestSelectStocks:
|
||||
"""close > MA5 且 MA5 > MA15 → 保留。"""
|
||||
# 构造 15 日 close 序列:上升 → close(末) > MA5 > MA15
|
||||
rising = [10.0 + i * 0.5 for i in range(15)] # 10→17
|
||||
df = _make_close_panel(["UP.XSHG"], [rising], end_date="2024-09-30", days=15)
|
||||
s = make_strategy(price_df_map={
|
||||
# 注意:_select_stocks 传 list,get_price 内部转 tuple → key 第一项必须是 tuple
|
||||
(("UP.XSHG",), ("close",), 15): df,
|
||||
panel = _make_close_wide(["UP.XSHG"], [rising], end_date="2024-09-30", days=15)
|
||||
s = make_strategy(panel_map={
|
||||
(("UP.XSHG",),): panel,
|
||||
})
|
||||
out = s._select_stocks(["UP.XSHG"], cur_date="2024-09-30")
|
||||
assert out == ["UP.XSHG"]
|
||||
@@ -230,9 +232,9 @@ class TestSelectStocks:
|
||||
def test_filter_close_below_ma_short(self):
|
||||
"""close < MA5 → 剔除(下行趋势)。"""
|
||||
falling = [20.0 - i * 0.5 for i in range(15)] # 20→13
|
||||
df = _make_close_panel(["DOWN.XSHG"], [falling], end_date="2024-09-30", days=15)
|
||||
s = make_strategy(price_df_map={
|
||||
(("DOWN.XSHG"), ("close",), 15): df,
|
||||
panel = _make_close_wide(["DOWN.XSHG"], [falling], end_date="2024-09-30", days=15)
|
||||
s = make_strategy(panel_map={
|
||||
(("DOWN.XSHG",),): panel,
|
||||
})
|
||||
out = s._select_stocks(["DOWN.XSHG"], cur_date="2024-09-30")
|
||||
assert out == []
|
||||
@@ -241,9 +243,9 @@ class TestSelectStocks:
|
||||
"""close > MA5 但 MA5 < MA15(下跌但末值小反弹)→ 剔除。"""
|
||||
# 前 10 日大涨(20→30),后 5 日跌(30→26):MA5 < MA15
|
||||
series = [20 + i for i in range(10)] + [30 - i for i in range(1, 6)] # 20..29, 29..25
|
||||
df = _make_close_panel(["FLAT.XSHG"], [series], end_date="2024-09-30", days=15)
|
||||
s = make_strategy(price_df_map={
|
||||
(("FLAT.XSHG"), ("close",), 15): df,
|
||||
panel = _make_close_wide(["FLAT.XSHG"], [series], end_date="2024-09-30", days=15)
|
||||
s = make_strategy(panel_map={
|
||||
(("FLAT.XSHG",),): panel,
|
||||
})
|
||||
out = s._select_stocks(["FLAT.XSHG"], cur_date="2024-09-30")
|
||||
# close=25, MA5 = mean(29,28,27,26,25)=27, MA15 = mean(all)=24.67
|
||||
@@ -252,17 +254,14 @@ class TestSelectStocks:
|
||||
|
||||
def test_insufficient_data_skipped(self):
|
||||
"""不足 ma_long=15 根 → 跳过。"""
|
||||
short_df = _make_close_panel(["NEW.XSHG"], [[10, 11, 12]], end_date="2024-09-30", days=15)
|
||||
s = make_strategy(price_df_map={
|
||||
(("NEW.XSHG"), ("close",), 15): short_df,
|
||||
# 直接给短宽表(3 行):panel.notna().sum()=3 < 15 → 全部剔除
|
||||
short_panel = pd.DataFrame(
|
||||
{"NEW.XSHG": [10.0, 11.0, 12.0]},
|
||||
index=pd.DatetimeIndex(["2024-09-28", "2024-09-29", "2024-09-30"]),
|
||||
)
|
||||
s = make_strategy(panel_map={
|
||||
(("NEW.XSHG",),): short_panel,
|
||||
})
|
||||
# 序列被 _make_close_panel 补齐到 15,这里改为真短数据
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = pd.DataFrame([
|
||||
{"time": pd.Timestamp("2024-09-28"), "code": "NEW.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-29"), "code": "NEW.XSHG", "close": 11.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "NEW.XSHG", "close": 12.0},
|
||||
])
|
||||
out = s._select_stocks(["NEW.XSHG"], cur_date="2024-09-30")
|
||||
assert out == []
|
||||
|
||||
@@ -278,9 +277,9 @@ class TestCalBuySign:
|
||||
# 上升序列:末值远高于均值
|
||||
idx_list = ["000300.XSHG", "000905.XSHG"]
|
||||
rising = [10.0 + i for i in range(30)] # 10→39
|
||||
df = _make_close_panel(idx_list, [rising, rising], end_date="2024-09-30", days=30)
|
||||
s = make_strategy(price_df_map={
|
||||
(tuple(idx_list), ("close",), 30): df,
|
||||
panel = _make_close_wide(idx_list, [rising, rising], end_date="2024-09-30", days=30)
|
||||
s = make_strategy(panel_map={
|
||||
(tuple(idx_list),): panel,
|
||||
})
|
||||
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is True
|
||||
|
||||
@@ -288,9 +287,9 @@ class TestCalBuySign:
|
||||
"""所有指数都跌破 30 日均线 → 占比 0% < 20% → 熊市(False)。"""
|
||||
idx_list = ["000300.XSHG", "000905.XSHG"]
|
||||
falling = [40.0 - i for i in range(30)] # 40→11
|
||||
df = _make_close_panel(idx_list, [falling, falling], end_date="2024-09-30", days=30)
|
||||
s = make_strategy(price_df_map={
|
||||
(tuple(idx_list), ("close",), 30): df,
|
||||
panel = _make_close_wide(idx_list, [falling, falling], end_date="2024-09-30", days=30)
|
||||
s = make_strategy(panel_map={
|
||||
(tuple(idx_list),): panel,
|
||||
})
|
||||
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is False
|
||||
|
||||
@@ -301,18 +300,18 @@ class TestCalBuySign:
|
||||
falling = [40.0 - i for i in range(30)]
|
||||
# 2 个 rising + 7 个 falling
|
||||
series_list = [rising, rising] + [falling] * 7
|
||||
df = _make_close_panel(idx_list, series_list, end_date="2024-09-30", days=30)
|
||||
s = make_strategy(price_df_map={
|
||||
(tuple(idx_list), ("close",), 30): df,
|
||||
panel = _make_close_wide(idx_list, series_list, end_date="2024-09-30", days=30)
|
||||
s = make_strategy(panel_map={
|
||||
(tuple(idx_list),): panel,
|
||||
})
|
||||
# 2/9 ≈ 0.222 > 0.2 → 牛市
|
||||
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is True
|
||||
|
||||
# 改为 1 个 rising:1/9 ≈ 0.111 < 0.2 → 熊市
|
||||
series_list_1 = [rising] + [falling] * 8
|
||||
df_1 = _make_close_panel(idx_list, series_list_1, end_date="2024-09-30", days=30)
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = df_1
|
||||
panel_1 = _make_close_wide(idx_list, series_list_1, end_date="2024-09-30", days=30)
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = panel_1
|
||||
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is False
|
||||
|
||||
|
||||
@@ -322,9 +321,9 @@ class TestHandleData:
|
||||
"""熊市信号 → 全部持仓清掉。"""
|
||||
cfg = MomentumTimingConfig(index_list=["IDX.XSHG"])
|
||||
s = make_strategy(config=cfg)
|
||||
# 触发熊市:get_price 返回下行 close
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = _make_close_panel(
|
||||
# 触发熊市:get_closes_panel 返回下行 close
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = _make_close_wide(
|
||||
["IDX.XSHG"], [[40.0 - i for i in range(30)]],
|
||||
end_date="2024-10-08", days=30,
|
||||
)
|
||||
@@ -356,35 +355,22 @@ class TestHandleData:
|
||||
rising_30 = [10.0 + i for i in range(30)] # 牛市信号用
|
||||
rising_15 = [10.0 + i for i in range(15)] # 均线筛选用
|
||||
|
||||
# 提供所有可能查询路径的 price 数据
|
||||
# 提供所有可能查询路径的 close 宽表
|
||||
idx_codes = ["IDX.XSHG"]
|
||||
stock_codes = ["CAND.XSHG"]
|
||||
|
||||
def _gp(security, **kwargs):
|
||||
fields = tuple(kwargs.get("fields") or [])
|
||||
# 1) _cal_buy_sign: idx_list, count=30
|
||||
if security == idx_codes and kwargs.get("count") == 30:
|
||||
return _make_close_panel(idx_codes, [rising_30], days=30)
|
||||
# 2) _cal_rps for index 股池:股票, start/end 模式
|
||||
if (
|
||||
isinstance(security, list)
|
||||
and security == stock_codes
|
||||
and kwargs.get("start_date")
|
||||
):
|
||||
return pd.DataFrame([
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "CAND.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-10-08"), "code": "CAND.XSHG", "close": 39.0},
|
||||
])
|
||||
# 3) _select_stocks: count=15
|
||||
if (
|
||||
isinstance(security, list)
|
||||
and security == stock_codes
|
||||
and kwargs.get("count") == cfg.ma_long
|
||||
):
|
||||
return _make_close_panel(stock_codes, [rising_15], days=15)
|
||||
return pd.DataFrame()
|
||||
def _gcp(symbols, start=None, end=None, interval="d", fq="raw"):
|
||||
if symbols == idx_codes:
|
||||
# _cal_buy_sign: idx_list, tail(30)
|
||||
return _make_close_wide(idx_codes, [rising_30], days=30)
|
||||
if isinstance(symbols, list) and symbols == stock_codes:
|
||||
# _cal_rps (2 行 start/end) 与 _select_stocks (15 行) 都用同一份上升序列
|
||||
# 取 15 行,_cal_rps 用首末(_shift_date 算出 start),
|
||||
# _select_stocks 用全部 15 行(ma_long=15)
|
||||
return _make_close_wide(stock_codes, [rising_15], days=15)
|
||||
return pd.DataFrame(index=pd.DatetimeIndex([]))
|
||||
|
||||
s.provider.get_price.side_effect = _gp
|
||||
s.provider.get_closes_panel.side_effect = _gcp
|
||||
|
||||
ctx = FakeContext(
|
||||
current_dt=datetime(2024, 10, 8, 9, 30),
|
||||
@@ -405,27 +391,27 @@ class TestHandleData:
|
||||
"""⚠️ 修复原始 bug 验证:handle_data 必须用 context.current_dt 计算 cur_date,
|
||||
不能用 datetime.date.today()(后者取真实今天)。
|
||||
"""
|
||||
# 用一个明显不同的 current_dt,确认 get_price 的 end_date 跟随它
|
||||
# 用一个明显不同的 current_dt,确认 get_closes_panel 的 end 跟随它
|
||||
cfg = MomentumTimingConfig(index_list=["IDX.XSHG"])
|
||||
s = make_strategy(config=cfg)
|
||||
captured_end_dates: List[Any] = []
|
||||
captured_ends: List[Any] = []
|
||||
|
||||
def _gp(security, **kwargs):
|
||||
# 记录 end_date 用于断言
|
||||
if kwargs.get("end_date"):
|
||||
captured_end_dates.append(str(kwargs["end_date"]))
|
||||
def _gcp(symbols, start=None, end=None, interval="d", fq="raw"):
|
||||
# 记录 end 用于断言
|
||||
if end is not None:
|
||||
captured_ends.append(str(end))
|
||||
# 下行 → 熊市(快速 return,不查其他)
|
||||
return _make_close_panel(
|
||||
return _make_close_wide(
|
||||
["IDX.XSHG"], [[40.0 - i for i in range(30)]],
|
||||
end_date=str(kwargs.get("end_date", "2024-10-08"))[:10],
|
||||
end_date=str(end or "2024-10-08")[:10],
|
||||
days=30,
|
||||
)
|
||||
|
||||
s.provider.get_price.side_effect = _gp
|
||||
s.provider.get_closes_panel.side_effect = _gcp
|
||||
ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30))
|
||||
s.handle_data(ctx)
|
||||
# 至少一次 get_price 的 end_date 是 "2024-10-08"(来自 current_dt),非今天
|
||||
assert any("2024-10-08" in d for d in captured_end_dates)
|
||||
# 至少一次 get_closes_panel 的 end 是 "2024-10-08"(来自 current_dt),非今天
|
||||
assert any("2024-10-08" in d for d in captured_ends)
|
||||
|
||||
|
||||
# =================== _find_stock_pool (取强舍弱) ===================
|
||||
@@ -441,30 +427,31 @@ class TestFindStockPool:
|
||||
config=cfg,
|
||||
)
|
||||
# 涨幅:A=+100%, B=+50%, C=0%, D=+30%, E=-10%
|
||||
rps_df_1 = pd.DataFrame([
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "A.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "A.XSHG", "close": 20.0},
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "B.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "B.XSHG", "close": 15.0},
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "C.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "C.XSHG", "close": 10.0},
|
||||
])
|
||||
rps_df_2 = pd.DataFrame([
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "D.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "D.XSHG", "close": 13.0},
|
||||
{"time": pd.Timestamp("2024-09-01"), "code": "E.XSHG", "close": 10.0},
|
||||
{"time": pd.Timestamp("2024-09-30"), "code": "E.XSHG", "close": 9.0},
|
||||
])
|
||||
rps_panel_1 = pd.DataFrame(
|
||||
{
|
||||
"A.XSHG": [10.0, 20.0],
|
||||
"B.XSHG": [10.0, 15.0],
|
||||
"C.XSHG": [10.0, 10.0],
|
||||
},
|
||||
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
|
||||
)
|
||||
rps_panel_2 = pd.DataFrame(
|
||||
{
|
||||
"D.XSHG": [10.0, 13.0],
|
||||
"E.XSHG": [10.0, 9.0],
|
||||
},
|
||||
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
|
||||
)
|
||||
|
||||
def _gp(security, **kwargs):
|
||||
if isinstance(security, list):
|
||||
if "A.XSHG" in security:
|
||||
return rps_df_1
|
||||
if "D.XSHG" in security:
|
||||
return rps_df_2
|
||||
return pd.DataFrame()
|
||||
def _gcp(symbols, start=None, end=None, interval="d", fq="raw"):
|
||||
if isinstance(symbols, list):
|
||||
if "A.XSHG" in symbols:
|
||||
return rps_panel_1
|
||||
if "D.XSHG" in symbols:
|
||||
return rps_panel_2
|
||||
return pd.DataFrame(index=pd.DatetimeIndex([]))
|
||||
|
||||
s.provider.get_price.side_effect = _gp
|
||||
s.provider.get_closes_panel.side_effect = _gcp
|
||||
out = s._find_stock_pool(
|
||||
["IDX1.XSHG", "IDX2.XSHG"], cur_date="2024-09-30", pre_date="2024-09-01",
|
||||
)
|
||||
|
||||
@@ -9,6 +9,10 @@
|
||||
- ✅ 5 日调仓周期:day_count % tc == 0 时选股+调仓,其他日 no-op
|
||||
- ✅ 等权 20 只
|
||||
- ❌ 对冲部分(已删,不测)
|
||||
|
||||
性能改造后:策略用 ``get_closes_panel`` 批量取 close 宽表(替代 get_price+逐只循环),
|
||||
**用 close.min/max 代理 low/high**(有意决策,见 small_cap._cal_momentum_score docstring)。
|
||||
测试 mock 同步切到 ``get_closes_panel`` 返回 close 宽表。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -33,14 +37,16 @@ def make_strategy(
|
||||
*,
|
||||
universe_stocks: Optional[List[str]] = None,
|
||||
fundamentals_df: Optional[pd.DataFrame] = None,
|
||||
price_df_map: Optional[Dict[Any, pd.DataFrame]] = None,
|
||||
panel_map: Optional[Dict[Any, pd.DataFrame]] = None,
|
||||
default_panel: Optional[pd.DataFrame] = None,
|
||||
config: Optional[SmallCapConfig] = None,
|
||||
) -> SmallCapStrategy:
|
||||
"""构造一个 mock provider + mock broker 驱动的策略。
|
||||
|
||||
- universe_stocks: get_index_stocks(universe, date) 返回的全市场候选列表
|
||||
- fundamentals_df: get_fundamentals_df 返回(index=code, cols=[market_cap, eps, ...])
|
||||
- price_df_map: get_price 按 (security_tuple, fields_tuple, count) 缓存的返回
|
||||
- panel_map: get_closes_panel 按 (symbols_tuple,) 缓存的宽表返回
|
||||
- default_panel: 所有未命中 panel_map 的查询返回的默认宽表(便于 _pick_stocks 流程测试)
|
||||
"""
|
||||
provider = MagicMock(name="provider")
|
||||
|
||||
@@ -73,24 +79,27 @@ def make_strategy(
|
||||
else:
|
||||
provider.get_fundamentals_df.return_value = pd.DataFrame()
|
||||
|
||||
# get_price 按 key 缓存
|
||||
def _normalize_key(k: Any) -> Any:
|
||||
# get_closes_panel 按 key 缓存:支持 (tuple(symbols),) 通配与精确 (tuple, start, end)
|
||||
def _normalize_syms(k: Any) -> Any:
|
||||
if isinstance(k, tuple) and k and isinstance(k[0], (list, tuple)):
|
||||
return (tuple(k[0]),) + tuple(k[1:])
|
||||
return k
|
||||
|
||||
price_df_map = {_normalize_key(k): v for k, v in (price_df_map or {}).items()}
|
||||
panel_map = {_normalize_syms(k): v for k, v in (panel_map or {}).items()}
|
||||
|
||||
def _get_price(security, **kwargs):
|
||||
sec_key = tuple(security) if isinstance(security, list) else security
|
||||
fields = tuple(kwargs.get("fields") or [])
|
||||
if kwargs.get("count") is not None:
|
||||
key = (sec_key, fields, kwargs.get("count"))
|
||||
else:
|
||||
key = (sec_key, fields, kwargs.get("start_date"), kwargs.get("end_date"))
|
||||
return price_df_map.get(key, pd.DataFrame())
|
||||
def _get_closes_panel(symbols, start=None, end=None, interval="d", fq="raw"):
|
||||
syms_key = tuple(symbols) if isinstance(symbols, list) else symbols
|
||||
exact = panel_map.get((syms_key, start, end))
|
||||
if exact is not None:
|
||||
return exact
|
||||
fallback = panel_map.get((syms_key,))
|
||||
if fallback is not None:
|
||||
return fallback
|
||||
if default_panel is not None:
|
||||
return default_panel
|
||||
return pd.DataFrame(index=pd.DatetimeIndex([]))
|
||||
|
||||
provider.get_price.side_effect = _get_price
|
||||
provider.get_closes_panel.side_effect = _get_closes_panel
|
||||
|
||||
broker = BrokerFacade()
|
||||
broker.order_target_value = MagicMock(return_value=MagicMock(filled=100))
|
||||
@@ -119,46 +128,29 @@ def _make_fundamentals_df(
|
||||
return df.set_index("code", drop=False)
|
||||
|
||||
|
||||
def _make_hlc_panel(
|
||||
def _make_close_wide(
|
||||
stocks: List[str],
|
||||
closes: List[List[float]],
|
||||
*,
|
||||
highs: Optional[List[List[float]]] = None,
|
||||
lows: Optional[List[List[float]]] = None,
|
||||
end_date: str = "2024-09-30",
|
||||
days: int = 130,
|
||||
) -> pd.DataFrame:
|
||||
"""构造 panel=False 风格的 close+high+low DataFrame。
|
||||
"""构造 get_closes_panel 风格的 close 宽表 (index=DatetimeIndex, columns=stocks)。
|
||||
|
||||
Args:
|
||||
stocks: 股票代码列表
|
||||
stocks: 股票代码列表(列名)
|
||||
closes: 每只股票的 close 序列(长度 <= days, 不足重复首值)
|
||||
highs: 同 close,None → 取 close
|
||||
lows: 同 close,None → 取 close
|
||||
days: 总 K 线根数(默认 130)
|
||||
"""
|
||||
end_dt = datetime.strptime(end_date, "%Y-%m-%d")
|
||||
dates = [
|
||||
(end_dt - timedelta(days=days - 1 - i)).strftime("%Y-%m-%d")
|
||||
for i in range(days)
|
||||
]
|
||||
rows = []
|
||||
for idx, code in enumerate(stocks):
|
||||
close_list = closes[idx]
|
||||
high_list = highs[idx] if highs else close_list
|
||||
low_list = lows[idx] if lows else close_list
|
||||
c_full = list(close_list) + [close_list[-1]] * (days - len(close_list))
|
||||
h_full = list(high_list) + [high_list[-1]] * (days - len(high_list))
|
||||
l_full = list(low_list) + [low_list[-1]] * (days - len(low_list))
|
||||
for d, c, h, l in zip(dates, c_full, h_full, l_full):
|
||||
rows.append({
|
||||
"time": pd.Timestamp(d),
|
||||
"code": code,
|
||||
"close": float(c),
|
||||
"high": float(h),
|
||||
"low": float(l),
|
||||
})
|
||||
return pd.DataFrame(rows)
|
||||
dates = pd.DatetimeIndex([
|
||||
end_dt - timedelta(days=days - 1 - i) for i in range(days)
|
||||
])
|
||||
data: Dict[str, List[float]] = {}
|
||||
for code, close_list in zip(stocks, closes):
|
||||
full = list(close_list) + [close_list[-1]] * (days - len(close_list))
|
||||
data[code] = [float(c) for c in full]
|
||||
return pd.DataFrame(data, index=dates)
|
||||
|
||||
|
||||
# =================== initialize ===================
|
||||
@@ -239,6 +231,7 @@ class TestCalMomentumScore:
|
||||
def test_score_formula_is_cur_minus_low_high_ma15(self):
|
||||
"""score = (cur-low_130) + (cur-high_130) + (cur-ma15)。
|
||||
|
||||
改造后用 close.min/max 代理 low/high(有意决策,见 small_cap docstring)。
|
||||
构造已知序列验证公式:
|
||||
- close 全 10(平):low=high=ma15=10,cur=10,score=0
|
||||
- close 上升:cur>low/high/ma15 → score 正
|
||||
@@ -248,13 +241,13 @@ class TestCalMomentumScore:
|
||||
rising = [10.0 + i * 0.1 for i in range(130)] # 10→22.9,cur=22.9
|
||||
falling = [23.0 - i * 0.1 for i in range(130)] # 23→10.1,cur=10.1
|
||||
|
||||
df = _make_hlc_panel(
|
||||
panel = _make_close_wide(
|
||||
["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"],
|
||||
[flat, rising, falling],
|
||||
end_date="2024-09-30", days=130,
|
||||
)
|
||||
s = make_strategy(price_df_map={
|
||||
(("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"), ("close", "high", "low"), 130): df,
|
||||
s = make_strategy(panel_map={
|
||||
(("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"),): panel,
|
||||
})
|
||||
out = s._cal_momentum_score(
|
||||
["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"], end_date="2024-09-30",
|
||||
@@ -262,11 +255,11 @@ class TestCalMomentumScore:
|
||||
|
||||
# FLAT: score = 0(全部相同)
|
||||
assert out.loc["FLAT.XSHG", "score"] == pytest.approx(0.0, abs=0.01)
|
||||
# UP: cur=22.9, low=10, high=22.9, ma15=mean([21.5..22.9])≈22.2
|
||||
# score = (22.9-10) + (22.9-22.9) + (22.9-22.2) = 12.9 + 0 + 0.7 ≈ 13.6
|
||||
# UP: 用 close 代理后 low=10(close min)、high=22.9(close max)、
|
||||
# ma15=mean([21.5..22.9])≈22.2,score = (22.9-10) + (22.9-22.9) + (22.9-22.2) ≈ 13.6
|
||||
assert out.loc["UP.XSHG", "score"] > 0
|
||||
# DOWN: cur=10.1, low=10.1, high=23, ma15≈10.8
|
||||
# score = (10.1-10.1) + (10.1-23) + (10.1-10.8) ≈ 0 + (-12.9) + (-0.7) ≈ -13.6
|
||||
# DOWN: 用 close 代理后 low=10.1、high=23、ma15≈10.8
|
||||
# score = (10.1-10.1) + (10.1-23) + (10.1-10.8) ≈ -13.6
|
||||
assert out.loc["DOWN.XSHG", "score"] < 0
|
||||
|
||||
def test_score_sorted_ascending(self):
|
||||
@@ -274,13 +267,13 @@ class TestCalMomentumScore:
|
||||
flat = [10.0] * 130
|
||||
rising = [10.0 + i * 0.1 for i in range(130)]
|
||||
falling = [23.0 - i * 0.1 for i in range(130)]
|
||||
df = _make_hlc_panel(
|
||||
panel = _make_close_wide(
|
||||
["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"],
|
||||
[flat, rising, falling],
|
||||
end_date="2024-09-30", days=130,
|
||||
)
|
||||
s = make_strategy(price_df_map={
|
||||
(("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"), ("close", "high", "low"), 130): df,
|
||||
s = make_strategy(panel_map={
|
||||
(("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"),): panel,
|
||||
})
|
||||
out = s._cal_momentum_score(
|
||||
["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"], end_date="2024-09-30",
|
||||
@@ -291,9 +284,9 @@ class TestCalMomentumScore:
|
||||
def test_insufficient_data_skipped(self):
|
||||
"""K 线序列不足/空 → 该股跳过(不在结果里)。"""
|
||||
s = make_strategy()
|
||||
# 让 provider.get_price 返回空 DataFrame
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = pd.DataFrame()
|
||||
# 让 provider.get_closes_panel 返回空 DataFrame
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = pd.DataFrame(index=pd.DatetimeIndex([]))
|
||||
out = s._cal_momentum_score(["EMPTY.XSHG"], end_date="2024-09-30")
|
||||
assert out.empty
|
||||
|
||||
@@ -320,14 +313,13 @@ class TestPickStocks:
|
||||
# 不传 price → _cal_momentum_score 会拿到空 df → 结果可能为空
|
||||
# 我们只验证 eps 过滤生效:在 fundamentals 过滤后 top_candidates 不含 B/C
|
||||
# 直接调 _pick_stocks 会因 price 空导致评分为空 → 返回空
|
||||
# 这里通过 mock price 给所有候选相同 close,看最终名单
|
||||
df = _make_hlc_panel(
|
||||
# 这里通过 mock get_closes_panel 给所有候选相同 close,看最终名单
|
||||
panel = _make_close_wide(
|
||||
["A.XSHG", "D.XSHG"], [[10.0] * 130, [10.0] * 130],
|
||||
end_date="2024-09-30", days=130,
|
||||
)
|
||||
# _pick_stocks 的 get_price 入参可能是 list 形式
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = df
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = panel
|
||||
ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30))
|
||||
out = s._pick_stocks(ctx)
|
||||
# eps>0 的 A/D 都进入候选,B/C 被剔
|
||||
@@ -350,13 +342,13 @@ class TestPickStocks:
|
||||
fundamentals_df=fund,
|
||||
config=cfg,
|
||||
)
|
||||
df = _make_hlc_panel(
|
||||
panel = _make_close_wide(
|
||||
["SMALL.XSHG", "MID.XSHG"],
|
||||
[[10.0] * 130, [10.0] * 130],
|
||||
end_date="2024-09-30", days=130,
|
||||
)
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = df
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = panel
|
||||
ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30))
|
||||
out = s._pick_stocks(ctx)
|
||||
# market_cap 升序后前 2 只 = SMALL/MID,BIG 被剔
|
||||
@@ -377,9 +369,9 @@ class TestPickStocks:
|
||||
)
|
||||
# 所有股票 close 相同 → score 相同 → 顺序由 sort_values stable 决定
|
||||
closes = [[10.0] * 130 for _ in stocks]
|
||||
df = _make_hlc_panel(stocks, closes, end_date="2024-09-30", days=130)
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = df
|
||||
panel = _make_close_wide(stocks, closes, end_date="2024-09-30", days=130)
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = panel
|
||||
ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30))
|
||||
out = s._pick_stocks(ctx)
|
||||
assert len(out) == 20
|
||||
@@ -404,12 +396,12 @@ class TestPickStocks:
|
||||
flat = [10.0] * 130
|
||||
rising = [10.0 + i * 0.1 for i in range(130)]
|
||||
falling = [23.0 - i * 0.1 for i in range(130)]
|
||||
df = _make_hlc_panel(
|
||||
panel = _make_close_wide(
|
||||
["DOWN.XSHG", "FLAT.XSHG", "UP.XSHG"],
|
||||
[falling, flat, rising], end_date="2024-09-30", days=130,
|
||||
)
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = df
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = panel
|
||||
ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30))
|
||||
out = s._pick_stocks(ctx)
|
||||
# 顺序:DOWN(score 最负) → FLAT(0),UP 被剔
|
||||
@@ -477,9 +469,9 @@ class TestHandleDataRebalance:
|
||||
fundamentals_df=fund,
|
||||
config=cfg,
|
||||
)
|
||||
df = _make_hlc_panel(["NEW.XSHG"], [[10.0] * 130], end_date="2024-09-30", days=130)
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = df
|
||||
df = _make_close_wide(["NEW.XSHG"], [[10.0] * 130], end_date="2024-09-30", days=130)
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = df
|
||||
ctx = FakeContext(
|
||||
current_dt=datetime(2024, 10, 8, 9, 30),
|
||||
positions={
|
||||
@@ -508,12 +500,12 @@ class TestHandleDataRebalance:
|
||||
fundamentals_df=fund,
|
||||
config=cfg,
|
||||
)
|
||||
df = _make_hlc_panel(
|
||||
df = _make_close_wide(
|
||||
["A.XSHG", "B.XSHG"], [[10.0] * 130, [10.0] * 130],
|
||||
end_date="2024-09-30", days=130,
|
||||
)
|
||||
s.provider.get_price.side_effect = None
|
||||
s.provider.get_price.return_value = df
|
||||
s.provider.get_closes_panel.side_effect = None
|
||||
s.provider.get_closes_panel.return_value = df
|
||||
ctx = FakeContext(
|
||||
current_dt=datetime(2024, 10, 8, 9, 30),
|
||||
positions={}, cash=1_000_000,
|
||||
|
||||
Reference in New Issue
Block a user