feat: 12个月窗口+资金面/板块维度 — 修预热期+网格相位bug,EM行业体系86板块全历史,v6.2达82.1%胜率/回撤9.68%

This commit is contained in:
hmo
2026-07-29 01:17:13 +08:00
parent 85ad6043d6
commit 0298aa7660
+130 -17
View File
@@ -88,7 +88,7 @@ STRATEGIES = {
"hh_only": True}},
"exit": {"tp_pct": 0.15, "sl_atr": 1.5, "max_hold_days": 20},
"sizing": {"kelly": True, "kelly_fraction": 0.5},
"eval_step": 5,
"eval_step": 1,
},
},
"v4.1": {
@@ -179,6 +179,22 @@ STRATEGIES.update({
"消融结果:单放量比+10笔保73%胜率/盈亏比2.92,单放ATR+4笔保78%胜率——两者是唯一不稀释质量的放宽,融合期望20+笔且保住70%胜率",
entry_overrides={"vol_ratio_min": 0.9, "vol_ratio_max": 2.0,
"atr_pct_min": 2.8, "atr_pct_max": 6.5}),
# E组: 资金面/板块增强
"v6.0": _v40_branch("v6.0", "资金流过滤",
"v4.0f + 资金5日净占比>-2.5(排除持续流出)",
"资金流归因:flow_5d<-2.7的持续流出组胜率仅37.5%,其余各桶55-78%。排除主力持续出逃的标的;flow_delta>4虽77.8%但会过度砍样本暂不启用",
entry_overrides={"vol_ratio_min": 0.9, "vol_ratio_max": 2.0,
"flow_5d_min": -2.5}),
"v6.1": _v40_branch("v6.1", "板块不追高",
"v4.0f + 板块MA20斜率≤1.0(排除已强涨板块)",
"板块归因:sector_slope>1.04的强涨板块入场胜率仅27.3%,而斜率≤0.08的回调/横盘板块胜率80%。与大盘/个股层面的'回调买'规律三层同构",
entry_overrides={"vol_ratio_min": 0.9, "vol_ratio_max": 2.0,
"sector_slope_max": 1.0}),
"v6.2": _v40_branch("v6.2", "资金+板块双滤",
"v4.0f + 资金5日净占比>-2.5 + 板块MA20斜率≤1.0",
"资金流(37.5%→55-78%)与板块(27.3%→80%)两个独立维度的负向排除叠加,期望在v4.0f基础上再提升胜率且不显著减样本",
entry_overrides={"vol_ratio_min": 0.9, "vol_ratio_max": 2.0,
"flow_5d_min": -2.5, "sector_slope_max": 1.0}),
})
@@ -216,33 +232,63 @@ def prepare_market_context(start_date, end_date):
}
def prepare_sector_context(start_date, end_date):
"""每日行业强度: 平均涨幅/净流入/当日排名 + 个股→行业映射"""
"""行业上下文: sector_index_daily(全历史) 提供板块趋势; sector_snapshots(近期) 补充净流入"""
global _SECTOR_CTX, _STOCK_SECTOR
_SECTOR_CTX, _STOCK_SECTOR = {}, {}
conn = sqlite3.connect(DB_PATH)
try:
rows = conn.execute("""
# 1. 板块指数历史(全周期)
try:
idx_rows = conn.execute("""
SELECT sector, date, close, change_pct FROM sector_index_daily
WHERE date >= ? AND date <= ? ORDER BY sector, date
""", (start_date, end_date)).fetchall()
except sqlite3.OperationalError:
idx_rows = []
# 每板块计算 MA20 和斜率
from collections import defaultdict
by_sector = defaultdict(list)
for sec, d, close, chg in idx_rows:
by_sector[sec].append((d, close, chg))
for sec, series in by_sector.items():
closes = [c for _, c, _ in series]
for i, (d, close, chg) in enumerate(series):
above = slope = None
if i >= 19:
ma20 = sum(closes[i-19:i+1]) / 20
above = close > ma20
if i >= 24:
ma20_5 = sum(closes[i-24:i-4]) / 20
if ma20_5 > 0:
slope = round((ma20 - ma20_5) / ma20_5 * 100, 3)
_SECTOR_CTX.setdefault(d, {})[sec] = {
'change': chg, 'above_ma20': above, 'slope': slope,
}
# 2. sector_snapshots 补充净流入和涨幅(近期,THS命名)
snap_rows = conn.execute("""
SELECT substr(m.timestamp,1,10) as d, s.name,
AVG(s.change_pct), SUM(s.net_inflow)
FROM sector_snapshots s JOIN market_snapshots m ON s.snapshot_id = m.id
WHERE m.timestamp >= ? AND m.timestamp <= ?
GROUP BY d, s.name
""", (start_date, end_date + ' 23:59')).fetchall()
# 优先 THS 源命名(与 sector_snapshots 同体系),证监会分类作兜底
_STOCK_SECTOR = {}
for d, name, chg, inflow in snap_rows:
e = _SECTOR_CTX.setdefault(d, {}).setdefault(name, {})
e['inflow'] = round(inflow or 0, 1)
if 'change' not in e or e.get('change') is None:
e['change'] = round(chg or 0, 2)
# 3. 个股→行业映射:优先 EM 体系(覆盖全,与 sector_index_daily 对齐),THS 兜底
try:
for code, sec in conn.execute("SELECT code, sector FROM stock_sectors_em").fetchall():
_STOCK_SECTOR[code] = sec
except sqlite3.OperationalError:
pass
for code, sec, src in conn.execute(
"SELECT code, sector_name, source FROM stock_sectors").fetchall():
if src == 'ths' or code not in _STOCK_SECTOR:
if code not in _STOCK_SECTOR and src == 'ths':
_STOCK_SECTOR[code] = sec
finally:
conn.close()
for d, name, chg, inflow in rows:
_SECTOR_CTX.setdefault(d, {})[name] = {'change': round(chg or 0, 2), 'inflow': round(inflow or 0, 1)}
for d in _SECTOR_CTX:
ranked = sorted(_SECTOR_CTX[d].items(), key=lambda x: -(x[1]['change']))
total = len(ranked)
for rank, (name, v) in enumerate(ranked):
v['rank_pct'] = round(rank / total, 3) if total else None # 0=最强
def mkt_ctx(date):
return _MKT_CTX.get(date, {})
@@ -254,6 +300,55 @@ def sector_ctx(code, date):
return _SECTOR_CTX.get(date, {}).get(sec, {})
# ══════════════════════════════════════════════════════
# 资金面上下文(stock_capital_flow 表)
# ══════════════════════════════════════════════════════
_FLOW_CTX = {} # code -> {date: main_pct}
_FLOW_SORTED = {} # code -> sorted list of (date, main_pct)
def prepare_flow_context(start_date, end_date):
"""加载个股主力资金净流入净占比历史"""
global _FLOW_CTX, _FLOW_SORTED
_FLOW_CTX, _FLOW_SORTED = {}, {}
conn = sqlite3.connect(DB_PATH)
try:
rows = conn.execute("""
SELECT code, date, main_pct FROM stock_capital_flow
WHERE date >= ? AND date <= ?
""", (start_date, end_date)).fetchall()
except sqlite3.OperationalError:
rows = [] # 表不存在时容忍
finally:
conn.close()
for code, d, pct in rows:
_FLOW_CTX.setdefault(code, {})[d] = pct
for code, dmap in _FLOW_CTX.items():
_FLOW_SORTED[code] = sorted(dmap.items())
def flow_ctx(code, date, bars_dates, idx):
"""资金流因子: 当日净占比 / 5日均值 / 5日趋势"""
series = _FLOW_SORTED.get(code)
if not series:
return {}
dmap = _FLOW_CTX[code]
# 找 date 之前(含)最近5个资金流数据点
dates = [d for d, _ in series if d <= date]
if not dates:
return {}
recent = dates[-5:]
prior = dates[-10:-5]
f = {'flow_pct': dmap.get(dates[-1])}
if recent:
vals = [dmap[d] for d in recent if dmap.get(d) is not None]
f['flow_5d'] = round(sum(vals) / len(vals), 2) if vals else None
if recent and prior:
v5 = [dmap[d] for d in recent if dmap.get(d) is not None]
p5 = [dmap[d] for d in prior if dmap.get(d) is not None]
if v5 and p5:
f['flow_delta'] = round(sum(v5)/len(v5) - sum(p5)/len(p5), 2)
return f
# ══════════════════════════════════════════════════════
# 入场过滤器
# ══════════════════════════════════════════════════════
@@ -287,6 +382,12 @@ def pass_filters(factors, filters):
# 行业
if not chk('sector_change', filters.get('sector_change_min'), filters.get('sector_change_max')): return False
if not chk('sector_rank_pct', None, filters.get('sector_rank_pct_max')): return False
if not chk('sector_slope', filters.get('sector_slope_min'), filters.get('sector_slope_max')): return False
if filters.get('sector_above_ma20') and factors.get('sector_above_ma20') is not True: return False
# 资金面
if not chk('flow_pct', filters.get('flow_pct_min'), filters.get('flow_pct_max')): return False
if not chk('flow_5d', filters.get('flow_5d_min'), filters.get('flow_5d_max')): return False
if not chk('flow_delta', filters.get('flow_delta_min'), filters.get('flow_delta_max')): return False
return True
@@ -345,8 +446,11 @@ def run_backtest(strategy_version, start_date, end_date, capital=1000000, save=T
filters = entry_cfg.get('filters', {})
step = cfg.get('eval_step', 5)
prepare_market_context(start_date, end_date)
# 指标预热:取数往前多取120天,保证窗口首日 MA60/RSI/ATR 等已收敛
fetch_start = (datetime.strptime(start_date, '%Y-%m-%d') - timedelta(days=120)).strftime('%Y-%m-%d')
prepare_market_context(fetch_start, end_date)
prepare_sector_context(start_date, end_date)
prepare_flow_context(fetch_start, end_date)
conn = sqlite3.connect(DB_PATH)
stocks = conn.execute("""
@@ -361,12 +465,16 @@ def run_backtest(strategy_version, start_date, end_date, capital=1000000, save=T
for code, name in stocks:
screened += 1
bars = _bars(code, start_date, end_date)
bars = _bars(code, fetch_start, end_date)
if not bars or len(bars) < 25:
continue
i = 20
while i < len(bars):
# 只在考察窗口内产生交易,预热期 bars 仅供指标计算
if bars[i].get('date', '') < start_date:
i += 1
continue
window = bars[:i+1]
sc = compute_single_score(window)
if sc is None:
@@ -390,6 +498,11 @@ def run_backtest(strategy_version, start_date, end_date, capital=1000000, save=T
factors['sector_change'] = sc_ctx.get('change')
factors['sector_rank_pct'] = sc_ctx.get('rank_pct')
factors['sector_inflow'] = sc_ctx.get('inflow')
factors['sector_above_ma20'] = sc_ctx.get('above_ma20')
factors['sector_slope'] = sc_ctx.get('slope')
# 资金面因子
fl = flow_ctx(code, date, None, i)
factors.update(fl)
if pass_filters(factors, filters):
atr_val = last.get('atr') or 0
@@ -516,9 +629,9 @@ def calc_summary(trades, capital):
# ══════════════════════════════════════════════════════
ANALYZE_FACTORS = ['rsi', 'adx', 'macd_hist', 'roc', 'atr_pct', 'dist_ma20', 'vol_ratio',
'ma20_slope', 'macd_hist_delta', 'rsi_delta', 'mkt_slope', 'mkt_roc',
'sector_change', 'sector_rank_pct', 'score']
'sector_change', 'sector_rank_pct', 'sector_slope', 'flow_pct', 'flow_5d', 'flow_delta', 'score']
BOOL_FACTORS = ['trend_aligned', 'hh_structure', 'hl_structure', 'adx_rising',
'mkt_above_ma20', 'near_high_20d']
'mkt_above_ma20', 'near_high_20d', 'sector_above_ma20']
def analyze_trades(strategy_version):
conn = sqlite3.connect(DB_PATH)