Add walk-forward backtest optimization to mitigate signal overfitting (item m)

Rolling train/test folds over 3 years of data auto-optimize the three
cheap-to-tune trading parameters (STRONG/BUY score thresholds, max hold
time) via grid search on each fold's train window, then evaluate purely
on the held-out test window. Stitching all out-of-sample results gives
an honest performance estimate uninflated by tuning against the same
data used to score it.

Split signal_scoring.py's expensive 13-algorithm scoring from its cheap
final threshold classification so grid search can replay many parameter
combinations without recomputing indicators each time. Moved the
backtest engine (fetch/precompute/simulate) out of the API layer into
app/services/backtest_engine.py so both /backtest/run and the new
walk-forward optimizer share one implementation instead of drifting
copies — same rationale as the earlier signal_service.py split (item h).

Also merges two long-diverged Alembic migration heads discovered while
adding the walk_forward_results table, so `alembic upgrade head` has a
single target again.

New: POST/GET/DELETE /walk-forward/* endpoints, a Walk-Forward tab on
the Backtest page (fold table, out-of-sample equity curve, run history).
19 new backend tests (153 total, all passing).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
Le
2026-07-04 09:14:32 +07:00
parent 95119b039e
commit 625c2b3773
12 changed files with 1932 additions and 368 deletions
+416
View File
@@ -0,0 +1,416 @@
"""Backtest engine — candle fetch, indicator precompute, and trade simulation.
Split out of `app/api/v1/backtest.py` (which now only holds the FastAPI
routes) so both the single-run `/backtest/run` endpoint and the
walk-forward optimizer (`app/services/walk_forward.py`) share one
implementation instead of two copies drifting apart.
"""
from __future__ import annotations
from collections import defaultdict
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from sqlalchemy import select, and_
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.candle import Candle
from app.models.symbol import Symbol
from app.models.exchange import Exchange
from app.services.indicator_service import (
bollinger_bands, rsi, sma, macd, supertrend,
volume_breakout, ichimoku, detect_divergence, market_structure,
)
from app.services.signal_scoring import (
_classify_signal_combined,
STRONG_BUY, BUY, STRONG_SELL, SELL,
)
# Min candles for warmup: BB(20) + RSI(14) + some room = 30
MIN_CANDLES = 30
async def _fetch_symbol(db: AsyncSession, symbol: str, exchange: str) -> Symbol | None:
"""Look up a Symbol row by (symbol, exchange) name."""
result = await db.execute(
select(Symbol)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(and_(Exchange.name == exchange, Symbol.symbol == symbol))
)
return result.scalar_one_or_none()
async def _fetch_candles(
db: AsyncSession,
symbol_id,
timeframe: str,
since: datetime,
until: datetime | None = None,
) -> list[Candle]:
"""Fetch candles for a symbol/timeframe in [since, until), ordered ascending."""
conditions = [
Candle.symbol_id == symbol_id,
Candle.timeframe == timeframe,
Candle.timestamp >= since,
]
if until is not None:
conditions.append(Candle.timestamp < until)
result = await db.execute(
select(Candle).where(and_(*conditions)).order_by(Candle.timestamp.asc())
)
return list(result.scalars().all())
def _precompute_indicators(candles: list[Candle], timeframe: str) -> dict:
"""Pre-compute all candle data & indicators ONCE for a candle range.
Split out so `walk_forward.py` can precompute indicators for a fold's
data once, then cheaply replay many threshold combinations against it
(see `_simulate_trades`).
"""
# MTF config
tf_minutes = {"15m": 15, "30m": 30, "1h": 60, "4h": 240}
main_minutes = tf_minutes.get(timeframe, 30)
mtf_config = []
for mtf_tf, mtf_minutes, mtf_w in [("15m", 15, 0.5), ("1h", 60, 1.5), ("4h", 240, 2.0)]:
if mtf_tf == timeframe:
continue
mult = mtf_minutes // main_minutes
if mult >= 1 and len(candles) >= mult * MIN_CANDLES:
mtf_config.append((mtf_tf, mult, mtf_w))
candle_dicts_full = [
{"high": float(c.high), "low": float(c.low),
"close": float(c.close), "open": float(c.open),
"volume": float(c.volume)}
for c in candles
]
close_prices_full = [float(c.close) for c in candles]
# Pre-compute indicators on full dataset (O(n) instead of O(n²))
bb_full = bollinger_bands(close_prices_full) or {}
rsi_full = rsi(close_prices_full) or []
sma_full = sma(close_prices_full, 20) or []
macd_full = macd(close_prices_full) or {}
st_full = supertrend(candle_dicts_full) or {}
vb_full = volume_breakout(candle_dicts_full) or []
ichi_full = ichimoku(candle_dicts_full) or {}
smc_full = market_structure(candle_dicts_full) or {}
# Pre-compute divergence ONCE (uses full arrays, indexes match)
rsi_div_full = detect_divergence(close_prices_full, rsi_full)
macd_hist_full = macd_full.get("histogram", []) if macd_full else []
macd_div_full = detect_divergence(close_prices_full, macd_hist_full)
# Pre-build MTF candles ONCE per MTF config
mtf_precomputed = []
for mtf_name, mtf_mult, mtf_w in mtf_config:
mtf_candles_list = []
for j in range(0, len(candle_dicts_full) - mtf_mult + 1, mtf_mult):
chunk = candle_dicts_full[j:j + mtf_mult]
mtf_candles_list.append({
"open": chunk[0]["open"],
"high": max(c["high"] for c in chunk),
"low": min(c["low"] for c in chunk),
"close": chunk[-1]["close"],
"volume": sum(c["volume"] for c in chunk),
})
if len(mtf_candles_list) >= MIN_CANDLES:
mtf_p = [c["close"] for c in mtf_candles_list]
mtf_precomputed.append({
"name": mtf_name,
"weight": mtf_w,
"mult": mtf_mult,
"candles_list": mtf_candles_list,
"close_prices": mtf_p,
"bb": bollinger_bands(mtf_p) or {},
"rsi": rsi(mtf_p) or [],
"sma": sma(mtf_p, 20) or [],
"macd": macd(mtf_p) or {},
"st": supertrend(mtf_candles_list) or {},
"vb": volume_breakout(mtf_candles_list) or [],
"ichi": ichimoku(mtf_candles_list) or {},
"smc": market_structure(mtf_candles_list) or {},
})
return {
"candle_dicts_full": candle_dicts_full,
"close_prices_full": close_prices_full,
"bb_full": bb_full,
"rsi_full": rsi_full,
"sma_full": sma_full,
"macd_full": macd_full,
"st_full": st_full,
"vb_full": vb_full,
"ichi_full": ichi_full,
"smc_full": smc_full,
"rsi_div_full": rsi_div_full,
"macd_div_full": macd_div_full,
"mtf_precomputed": mtf_precomputed,
}
def _simulate_trades(
candles: list[Candle],
precomputed: dict,
trade_size: Decimal = Decimal("10"),
strong_threshold: float = 4.0,
signal_threshold: float = 1.0,
max_hold_candles: int = 48,
active_from_index: int = MIN_CANDLES,
) -> tuple[list[dict], list[dict]]:
"""Replay signal classification + trade simulation over precomputed indicators.
`active_from_index` lets callers pass extra warmup candles before the
window they actually want simulated (e.g. walk-forward fold
boundaries) — candles before this index are used only so indicators
have enough lookback, never turned into signals/trades.
"""
close_prices_full = precomputed["close_prices_full"]
bb_full = precomputed["bb_full"]
rsi_full = precomputed["rsi_full"]
sma_full = precomputed["sma_full"]
macd_full = precomputed["macd_full"]
st_full = precomputed["st_full"]
vb_full = precomputed["vb_full"]
ichi_full = precomputed["ichi_full"]
smc_full = precomputed["smc_full"]
rsi_div_full = precomputed["rsi_div_full"]
macd_div_full = precomputed["macd_div_full"]
mtf_precomputed = precomputed["mtf_precomputed"]
all_signals = []
trades = []
current_position = None
start_index = max(MIN_CANDLES, active_from_index)
for i in range(start_index, len(candles)):
candle = candles[i]
latest_close = close_prices_full[i]
timestamp = candle.timestamp.isoformat()
# Slice pre-computed arrays (O(i) but ~100x faster than recomputing)
clip = i + 1
def _safe_slice(v):
return v[:clip] if v is not None and hasattr(v, '__getitem__') else v
bb_data = {k: _safe_slice(v) for k, v in bb_full.items()} if bb_full else {}
rsi_data = rsi_full[:clip] if rsi_full else []
sma_data = sma_full[:clip] if sma_full else []
macd_data = {k: _safe_slice(v) for k, v in macd_full.items()} if macd_full else {}
st_data = {k: _safe_slice(v) for k, v in st_full.items()} if st_full else {}
vb_data = vb_full[:clip] if vb_full else []
ichi_data = {k: _safe_slice(v) for k, v in ichi_full.items()} if ichi_full else {}
smc_data = {k: _safe_slice(v) for k, v in smc_full.items()} if smc_full else {}
# MTF votes — use precomputed MTF indicators, sliced to current MTF candle index
mtf_votes = []
for mtf in mtf_precomputed:
# Which MTF candle corresponds to main candle i?
mtf_idx = i // mtf["mult"]
if mtf_idx < MIN_CANDLES or mtf_idx >= len(mtf["close_prices"]):
continue
clip_mtf = mtf_idx + 1
def _safe_slice_mtf(v):
return v[:clip_mtf] if v is not None and hasattr(v, '__getitem__') else v
mtf_s, *_ = _classify_signal_combined(
mtf["close_prices"][mtf_idx],
{k: _safe_slice_mtf(v) for k, v in mtf["bb"].items()},
mtf["rsi"][:clip_mtf],
mtf["sma"][:clip_mtf],
{k: _safe_slice_mtf(v) for k, v in mtf["macd"].items()} if mtf["macd"] else None,
{k: _safe_slice_mtf(v) for k, v in mtf["st"].items()} if mtf["st"] else None,
mtf["vb"][:clip_mtf] if mtf["vb"] else None,
{k: _safe_slice_mtf(v) for k, v in mtf["ichi"].items()} if mtf["ichi"] else None,
(None, None), (None, None),
{k: _safe_slice_mtf(v) for k, v in mtf["smc"].items()} if mtf["smc"] else None,
)
if mtf_s:
mtf_votes.append((mtf_s, "", mtf["weight"]))
signal_type, strength, *_ = _classify_signal_combined(
latest_close, bb_data, rsi_data, sma_data,
macd_data, st_data, vb_data, ichi_data,
rsi_div_full, macd_div_full, smc_data, mtf_votes or None,
strong_threshold=strong_threshold, signal_threshold=signal_threshold,
)
if signal_type:
all_signals.append({
"time": timestamp, "signal": signal_type,
"strength": strength or "", "price": latest_close,
})
# PnL simulation
if signal_type in (STRONG_BUY, BUY):
if current_position and current_position["direction"] == "SHORT":
if signal_type == STRONG_BUY:
entry = current_position["entry_price"]
qty = current_position["quantity"]
pnl = (entry - latest_close) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
})
trades.append(current_position)
current_position = None
else:
continue
if not current_position:
qty = float(trade_size) / latest_close
current_position = {
"direction": "LONG", "entry_price": latest_close,
"entry_time": timestamp, "quantity": qty,
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
}
elif signal_type in (STRONG_SELL, SELL):
if current_position and current_position["direction"] == "LONG":
if signal_type == STRONG_SELL:
entry = current_position["entry_price"]
qty = current_position["quantity"]
pnl = (latest_close - entry) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
})
trades.append(current_position)
current_position = None
else:
continue
if not current_position:
qty = float(trade_size) / latest_close
current_position = {
"direction": "SHORT", "entry_price": latest_close,
"entry_time": timestamp, "quantity": qty,
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
}
# Time limit
if current_position and current_position["status"] == "OPEN":
hold = i - current_position["entry_index"]
if hold >= max_hold_candles:
entry = current_position["entry_price"]
qty = current_position["quantity"]
if current_position["direction"] == "LONG":
pnl = (latest_close - entry) * qty
else:
pnl = (entry - latest_close) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "TIME_LIMIT",
})
trades.append(current_position)
current_position = None
# Close final position
if current_position and current_position["status"] == "OPEN":
last_close = float(candles[-1].close)
entry = current_position["entry_price"]
qty = current_position["quantity"]
if current_position["direction"] == "LONG":
pnl = (last_close - entry) * qty
else:
pnl = (entry - last_close) * qty
current_position.update({
"exit_price": last_close,
"exit_time": candles[-1].timestamp.isoformat(),
"pnl": pnl, "status": "CLOSED", "exit_reason": "END_OF_DATA",
})
trades.append(current_position)
return all_signals, trades
def _compute_stats(all_signals: list[dict], trades: list[dict]) -> dict:
"""Reduce raw signals/trades into the summary stats block used by
both the single-run backtest and each walk-forward fold."""
counts = defaultdict(int)
for s in all_signals:
counts[s["signal"]] += 1
closed_trades = [t for t in trades if t.get("status") == "CLOSED"]
winning_trades = [t for t in closed_trades if t.get("pnl", 0) > 0]
losing_trades = [t for t in closed_trades if t.get("pnl", 0) <= 0]
total_pnl = sum(t.get("pnl", 0) for t in closed_trades)
gross_profit = sum(t.get("pnl", 0) for t in winning_trades)
gross_loss = sum(t.get("pnl", 0) for t in losing_trades)
win_rate = round(len(winning_trades) / len(closed_trades) * 100, 1) if closed_trades else 0
profit_factor = round(abs(gross_profit / gross_loss), 2) if gross_loss != 0 else None
avg_win = round(gross_profit / len(winning_trades), 2) if winning_trades else None
avg_loss = round(gross_loss / len(losing_trades), 2) if losing_trades else None
best_trade = max(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
worst_trade = min(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
return {
"signal_counts": dict(counts),
"total_signals": len(all_signals),
"recent_signals": all_signals[-15:],
"trades": {
"total": len(closed_trades),
"wins": len(winning_trades),
"losses": len(losing_trades),
"win_rate": win_rate,
"total_pnl": round(total_pnl, 2),
"profit_factor": profit_factor,
"avg_win": avg_win,
"avg_loss": avg_loss,
"best_trade": {
"direction": best_trade.get("direction"),
"entry_price": round(best_trade["entry_price"], 4),
"exit_price": round(best_trade["exit_price"], 4),
"pnl": round(best_trade["pnl"], 2),
"entry_signal": best_trade.get("entry_signal"),
} if best_trade else None,
"worst_trade": {
"direction": worst_trade.get("direction"),
"entry_price": round(worst_trade["entry_price"], 4),
"exit_price": round(worst_trade["exit_price"], 4),
"pnl": round(worst_trade["pnl"], 2),
"entry_signal": worst_trade.get("entry_signal"),
} if worst_trade else None,
"per_signal": {},
"recent": closed_trades[-10:],
},
}
async def run_backtest(
db: AsyncSession,
symbol: str,
exchange: str,
timeframe: str = "30m",
days: int = 7,
trade_size: Decimal = Decimal("10"),
) -> dict:
"""Run a single backtest over the last `days` days and return structured results."""
db_symbol = await _fetch_symbol(db, symbol, exchange)
if not db_symbol:
return {"error": f"Symbol {symbol} not found on {exchange}"}
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
candles = await _fetch_candles(db, db_symbol.id, timeframe, since=cutoff)
if len(candles) < MIN_CANDLES:
if candles:
span_hours = (candles[-1].timestamp - candles[0].timestamp).total_seconds() / 3600
if span_hours >= 24:
avail = f"{span_hours/24:.0f}d"
else:
avail = f"{span_hours:.0f}h"
return {"error": f"Only ~{avail} data available (need at least {MIN_CANDLES} candles for {timeframe}). Try more days or a higher timeframe (4h)."}
return {"error": f"No candle data found for {symbol} on {timeframe}. The exchange may not support this pair."}
precomputed = _precompute_indicators(candles, timeframe)
all_signals, trades = _simulate_trades(candles, precomputed, trade_size)
stats = _compute_stats(all_signals, trades)
return {
"symbol": symbol,
"exchange": exchange,
"timeframe": timeframe,
"days": days,
"candles_count": len(candles),
**stats,
}
+105 -38
View File
@@ -165,7 +165,7 @@ def _classify_signal_bb(
return None, None
def _classify_signal_combined(
def _compute_adjusted_score(
close_price: float,
bb: dict[str, list[float]],
rsi: list[float] | None,
@@ -185,47 +185,37 @@ def _classify_signal_combined(
candlestick_score: float | None = None,
rates: dict[str, float] | None = None,
enabled_strategies: list[str] | None = None,
) -> tuple[Optional[str], Optional[str], float, dict[str, float]]:
"""Classify market state using 13-algorithm voting with win-rate boosting.
) -> tuple[Optional[str], Optional[str], float, float, dict[str, float]]:
"""Run the 13-algorithm vote and reduce it to a single adjusted score.
Algorithms:
1. Double BB + RSI
2. MACD Crossover
3. SuperTrend
4. Volume Breakout
5. Ichimoku Cloud
6. Divergence Detection (RSI + MACD)
7. 🌤️ Market Structure (SMC) — BOS, CHoCH, OB
8. 🔄 Multi-Timeframe (15m + 1h + 4h)
9. 📊 OBV (On-Balance Volume) Crossover
10. 🔄 Stochastic RSI Crossover
11. 💰 MFI (Money Flow Index)
12. 🕯️ FVG (Fair Value Gap)
13. 🕯️ Candlestick Patterns (30+ patterns)
This is the expensive, threshold-independent half of signal
classification — algorithms 1-13, win-rate boosting, correlation
dampening, and dynamic normalization. It does NOT decide the final
signal type; that is a cheap final step in `_classify_signal_combined`
(or `_score_to_signal`) so callers that need to try many threshold
combinations (e.g. walk-forward parameter search) can compute this
once per candle and replay different thresholds against it cheaply.
Each algorithm votes: BUY (+1/+2), SELL (-1/-2), or NEUTRAL (0).
If *rates* is provided, each strategy's raw score is boosted by its
historical win rate before the final classification.
Returns (signal_type, strength, confidence, raw_scores) where
confidence is a 0-1 float and raw_scores is a dict of all 9
algorithm scores for ML feature collection.
Returns (override_signal, override_strength, adjusted_score, confidence,
raw_scores). When override_signal is not None (SQUEEZE_ALERT,
CAUTION_LONG, CAUTION_SHORT), the caller must return it as-is —
it bypasses threshold-based classification entirely.
"""
# ── NaN/Inf guard: reject any invalid price before processing ──
if not math.isfinite(close_price) or close_price <= 0:
logger.warning("_classify_signal_combined: invalid close_price=%s, returning NEUTRAL", close_price)
return None, None, 0.0, {}
return None, None, 0.0, 0.0, {}
# ── Special signals (override) ──
squeeze = _detect_squeeze(bb)
if squeeze:
return SQUEEZE_ALERT, "MODERATE", 0.5, {}
return SQUEEZE_ALERT, "MODERATE", 0.0, 0.5, {}
# P2-2: Call _classify_signal_bb ONCE, reuse result for both
# early-return check AND the raw_scores vote
bb_type, bb_strength = _classify_signal_bb(close_price, bb, rsi, sma)
if bb_type in (CAUTION_LONG, CAUTION_SHORT):
return bb_type, "MODERATE", 0.5, {}
return bb_type, "MODERATE", 0.0, 0.5, {}
# ── Collect per-strategy raw scores ──
raw_scores: dict[str, float] = {
@@ -504,18 +494,95 @@ def _classify_signal_combined(
else:
adjusted_score = total_score
# ── Final classification from boosted score ──
# 🔧 Dynamic thresholds: STRONG needs effective 4.0, BUY/SELL needs 1.0
if adjusted_score >= 4.0:
return STRONG_BUY, "STRONG", confidence, raw_scores
elif adjusted_score >= 1.0:
return BUY, "MODERATE", confidence, raw_scores
elif adjusted_score <= -4.0:
return STRONG_SELL, "STRONG", confidence, raw_scores
elif adjusted_score <= -1.0:
return SELL, "MODERATE", confidence, raw_scores
return None, None, adjusted_score, confidence, raw_scores
return None, None, confidence, raw_scores
def _score_to_signal(
adjusted_score: float,
strong_threshold: float = 4.0,
signal_threshold: float = 1.0,
) -> tuple[Optional[str], Optional[str]]:
"""Turn an adjusted score into a signal type — the cheap, threshold-only
half of classification. Split out from `_compute_adjusted_score` so
walk-forward parameter search can replay many threshold combinations
against an already-computed score array without re-running the 13
algorithms each time.
🔧 Dynamic thresholds: STRONG needs effective `strong_threshold` (default
4.0), BUY/SELL needs `signal_threshold` (default 1.0).
"""
if adjusted_score >= strong_threshold:
return STRONG_BUY, "STRONG"
elif adjusted_score >= signal_threshold:
return BUY, "MODERATE"
elif adjusted_score <= -strong_threshold:
return STRONG_SELL, "STRONG"
elif adjusted_score <= -signal_threshold:
return SELL, "MODERATE"
return None, None
def _classify_signal_combined(
close_price: float,
bb: dict[str, list[float]],
rsi: list[float] | None,
sma: list[float] | None,
macd_data: dict | None,
st_data: dict | None,
vol_data: list | None,
ichi_data: dict | None = None,
rsi_div: tuple = (None, None),
macd_div: tuple = (None, None),
smc_data: dict | None = None,
mtf_votes: list[tuple[Optional[str], Optional[str], float]] | None = None,
obv_data: list | None = None,
stoch_rsi_data: dict | None = None,
mfi_data: list | None = None,
fvg_data: dict | None = None,
candlestick_score: float | None = None,
rates: dict[str, float] | None = None,
enabled_strategies: list[str] | None = None,
strong_threshold: float = 4.0,
signal_threshold: float = 1.0,
) -> tuple[Optional[str], Optional[str], float, dict[str, float]]:
"""Classify market state using 13-algorithm voting with win-rate boosting.
Algorithms:
1. Double BB + RSI
2. MACD Crossover
3. SuperTrend
4. Volume Breakout
5. Ichimoku Cloud
6. Divergence Detection (RSI + MACD)
7. 🌤️ Market Structure (SMC) — BOS, CHoCH, OB
8. 🔄 Multi-Timeframe (15m + 1h + 4h)
9. 📊 OBV (On-Balance Volume) Crossover
10. 🔄 Stochastic RSI Crossover
11. 💰 MFI (Money Flow Index)
12. 🕯️ FVG (Fair Value Gap)
13. 🕯️ Candlestick Patterns (30+ patterns)
Each algorithm votes: BUY (+1/+2), SELL (-1/-2), or NEUTRAL (0).
If *rates* is provided, each strategy's raw score is boosted by its
historical win rate before the final classification. *strong_threshold*
and *signal_threshold* control the final cutoffs (see `_score_to_signal`)
— left at their defaults for live trading; walk-forward backtesting
overrides them during parameter search.
Returns (signal_type, strength, confidence, raw_scores) where
confidence is a 0-1 float and raw_scores is a dict of all 9
algorithm scores for ML feature collection.
"""
override_signal, override_strength, adjusted_score, confidence, raw_scores = _compute_adjusted_score(
close_price, bb, rsi, sma, macd_data, st_data, vol_data, ichi_data,
rsi_div, macd_div, smc_data, mtf_votes, obv_data, stoch_rsi_data,
mfi_data, fvg_data, candlestick_score, rates, enabled_strategies,
)
if override_signal is not None:
return override_signal, override_strength, confidence, raw_scores
signal_type, strength = _score_to_signal(adjusted_score, strong_threshold, signal_threshold)
return signal_type, strength, confidence, raw_scores
def _calculate_pnl(
+303
View File
@@ -0,0 +1,303 @@
"""Walk-forward backtest optimization.
Splits historical data into rolling train/test folds, auto-optimizes a
small parameter grid on each fold's train window, then evaluates the
optimized parameters on that fold's held-out test window
(out-of-sample). Stitching all out-of-sample test results together
gives an honest performance estimate that isn't inflated by tuning
parameters against the same data used to score them — see item (m) in
theo_doi_trading-portal_v6.md.
Only the three cheaply-tunable "when to enter/exit" parameters are
optimized (see `_score_to_signal` in signal_scoring.py): the two score
thresholds and the max hold time. The 13-algorithm voting internals
(RSI/MFI/etc. cutoffs) are not parameterized — doing so would require a
much larger, riskier refactor of signal_scoring.py.
"""
from __future__ import annotations
import math
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from itertools import product
from sqlalchemy.ext.asyncio import AsyncSession
from app.services.backtest_engine import (
MIN_CANDLES,
_fetch_symbol,
_fetch_candles,
_precompute_indicators,
_simulate_trades,
_compute_stats,
)
# Parameters that are cheap to grid-search: they only affect the final
# threshold/exit logic, not the 13-algorithm scoring itself, so replaying
# them against already-precomputed indicators is fast (see _simulate_trades).
DEFAULT_PARAM_GRID: dict[str, list[float]] = {
"strong_threshold": [3.5, 4.0, 4.5],
"signal_threshold": [0.75, 1.0, 1.5],
"max_hold_candles": [24, 48, 96],
}
MIN_TRADES_PER_FOLD = 5 # reject param combos too sparse to trust
WARMUP_BUFFER_CANDLES = 60 # extra history fetched before each window so indicators aren't cold at window start
_TF_MINUTES = {"15m": 15, "30m": 30, "1h": 60, "4h": 240, "1d": 1440}
def _fold_score(closed_trades: list[dict], min_trades: int = MIN_TRADES_PER_FOLD) -> float:
"""Per-trade Sharpe-like objective, scaled by sqrt(n) trades.
Favors consistent edge over one lucky trade, and rejects parameter
combinations with too few trades to be statistically meaningful
(returns -inf so they never win a grid search).
"""
if len(closed_trades) < min_trades:
return float("-inf")
pnls = [float(t.get("pnl", 0.0)) for t in closed_trades]
n = len(pnls)
mean = sum(pnls) / n
variance = sum((p - mean) ** 2 for p in pnls) / n
std = math.sqrt(variance)
if std == 0:
return mean * math.sqrt(n)
return (mean / std) * math.sqrt(n)
def generate_folds(
total_days: int,
train_days: int,
test_days: int,
anchor: datetime | None = None,
) -> list[dict]:
"""Generate rolling folds: a fixed-size train window sliding forward
by `test_days` each step, anchored to `anchor` (default: now) counting
back `total_days`. Each fold covers [train_start, train_end) train +
[train_end, test_end) test, walked forward in chronological order.
"""
anchor = anchor or datetime.now(timezone.utc)
origin = anchor - timedelta(days=total_days)
folds = []
offset = 0
while True:
train_start = origin + timedelta(days=offset)
train_end = train_start + timedelta(days=train_days)
test_end = train_end + timedelta(days=test_days)
if test_end > anchor:
break
folds.append({
"fold_index": len(folds),
"train_start": train_start,
"train_end": train_end,
"test_start": train_end,
"test_end": test_end,
})
offset += test_days
return folds
def _warmup_days(timeframe: str, buffer_candles: int = WARMUP_BUFFER_CANDLES) -> int:
minutes = _TF_MINUTES.get(timeframe, 30)
return max(1, math.ceil(buffer_candles * minutes / 1440))
async def _prepare_window(
db: AsyncSession,
symbol_id,
timeframe: str,
window_start: datetime,
window_end: datetime,
) -> tuple[list, dict, int] | None:
"""Fetch candles for [window_start - warmup, window_end) and precompute
indicators. Returns (candles, precomputed, active_from_index), where
active_from_index is the candle index at which window_start begins —
candles before it exist only to warm up indicators, and are never
turned into signals/trades. Returns None if there isn't enough data.
"""
warmup_start = window_start - timedelta(days=_warmup_days(timeframe))
candles = await _fetch_candles(db, symbol_id, timeframe, since=warmup_start, until=window_end)
if len(candles) < MIN_CANDLES:
return None
active_from_index = None
for idx, c in enumerate(candles):
if c.timestamp >= window_start:
active_from_index = idx
break
if active_from_index is None:
return None # no candles actually within the window itself
precomputed = _precompute_indicators(candles, timeframe)
return candles, precomputed, active_from_index
def _run_combo(candles, precomputed, trade_size, params, active_from_index):
return _simulate_trades(
candles, precomputed, trade_size,
strong_threshold=params["strong_threshold"],
signal_threshold=params["signal_threshold"],
max_hold_candles=int(params["max_hold_candles"]),
active_from_index=active_from_index,
)
def _grid_search(
candles: list,
precomputed: dict,
param_grid: dict[str, list[float]],
trade_size: Decimal,
active_from_index: int,
) -> tuple[dict[str, float], float, dict]:
"""Try every combination in param_grid, return the one that scores best
on the train window by `_fold_score`."""
keys = list(param_grid.keys())
best_params: dict[str, float] | None = None
best_score = float("-inf")
best_stats: dict = {}
for combo in product(*(param_grid[k] for k in keys)):
params = dict(zip(keys, combo))
all_signals, trades = _run_combo(candles, precomputed, trade_size, params, active_from_index)
closed_trades = [t for t in trades if t.get("status") == "CLOSED"]
score = _fold_score(closed_trades)
if score > best_score:
best_score = score
best_params = params
best_stats = _compute_stats(all_signals, trades)
if best_params is None:
# Every combo scored -inf (too few trades) — still report the
# grid's first combination so the fold has *something* to show.
best_params = {k: param_grid[k][0] for k in keys}
all_signals, trades = _run_combo(candles, precomputed, trade_size, best_params, active_from_index)
best_stats = _compute_stats(all_signals, trades)
return best_params, best_score, best_stats
def _max_drawdown_pct(equity_curve: list[float]) -> float:
"""Max peak-to-trough decline of a cumulative-PnL equity curve, as a
percentage of the running peak."""
if not equity_curve:
return 0.0
peak = equity_curve[0]
max_dd = 0.0
for v in equity_curve:
peak = max(peak, v)
if peak > 0:
max_dd = max(max_dd, (peak - v) / peak * 100)
return round(max_dd, 2)
async def run_walk_forward(
db: AsyncSession,
symbol: str,
exchange: str,
timeframe: str = "4h",
total_days: int = 1095,
train_days: int = 270,
test_days: int = 90,
trade_size: Decimal = Decimal("10"),
param_grid: dict[str, list[float]] | None = None,
) -> dict:
"""Run a full walk-forward analysis: optimize params per fold on the
train window, evaluate out-of-sample on the test window, then stitch
all out-of-sample results into one honest performance estimate."""
param_grid = param_grid or DEFAULT_PARAM_GRID
db_symbol = await _fetch_symbol(db, symbol, exchange)
if not db_symbol:
return {"error": f"Symbol {symbol} not found on {exchange}"}
folds_spec = generate_folds(total_days, train_days, test_days)
if not folds_spec:
return {"error": f"total_days ({total_days}) too small for train_days+test_days ({train_days}+{test_days})"}
fold_results = []
stitched_oos_trades: list[dict] = []
for spec in folds_spec:
train_window = await _prepare_window(db, db_symbol.id, timeframe, spec["train_start"], spec["train_end"])
if train_window is None:
continue
train_candles, train_precomputed, train_active_from = train_window
best_params, _train_score, train_stats = _grid_search(
train_candles, train_precomputed, param_grid, trade_size, train_active_from,
)
test_window = await _prepare_window(db, db_symbol.id, timeframe, spec["test_start"], spec["test_end"])
if test_window is None:
continue
test_candles, test_precomputed, test_active_from = test_window
test_signals, test_trades = _run_combo(test_candles, test_precomputed, trade_size, best_params, test_active_from)
test_stats = _compute_stats(test_signals, test_trades)
stitched_oos_trades.extend(t for t in test_trades if t.get("status") == "CLOSED")
fold_results.append({
"fold_index": spec["fold_index"],
"train_start": spec["train_start"].isoformat(),
"train_end": spec["train_end"].isoformat(),
"test_start": spec["test_start"].isoformat(),
"test_end": spec["test_end"].isoformat(),
"best_params": best_params,
"in_sample": {
"trades": train_stats["trades"]["total"],
"win_rate": train_stats["trades"]["win_rate"],
"total_pnl": train_stats["trades"]["total_pnl"],
"profit_factor": train_stats["trades"]["profit_factor"],
},
"out_of_sample": {
"trades": test_stats["trades"]["total"],
"win_rate": test_stats["trades"]["win_rate"],
"total_pnl": test_stats["trades"]["total_pnl"],
"profit_factor": test_stats["trades"]["profit_factor"],
},
})
if not fold_results:
return {"error": "No fold had enough candle data to run — try a larger total_days or a lower timeframe."}
# Stitch OOS trades chronologically — this is the walk-forward's
# headline number, the only one that hasn't seen the data it's
# evaluated on.
stitched_oos_trades.sort(key=lambda t: t["entry_time"])
oos_wins = [t for t in stitched_oos_trades if t.get("pnl", 0) > 0]
oos_losses = [t for t in stitched_oos_trades if t.get("pnl", 0) <= 0]
oos_gross_profit = sum(t.get("pnl", 0) for t in oos_wins)
oos_gross_loss = sum(t.get("pnl", 0) for t in oos_losses)
oos_total_pnl = sum(t.get("pnl", 0) for t in stitched_oos_trades)
oos_win_rate = round(len(oos_wins) / len(stitched_oos_trades) * 100, 1) if stitched_oos_trades else 0.0
oos_profit_factor = round(abs(oos_gross_profit / oos_gross_loss), 2) if oos_gross_loss != 0 else None
equity_curve = [0.0]
running = 0.0
for t in stitched_oos_trades:
running += float(t.get("pnl", 0))
equity_curve.append(round(running, 4))
return {
"symbol": symbol,
"exchange": exchange,
"timeframe": timeframe,
"total_days": total_days,
"train_days": train_days,
"test_days": test_days,
"param_grid": param_grid,
"folds": fold_results,
"out_of_sample_summary": {
"trades": len(stitched_oos_trades),
"wins": len(oos_wins),
"losses": len(oos_losses),
"win_rate": oos_win_rate,
"total_pnl": round(oos_total_pnl, 2),
"profit_factor": oos_profit_factor,
"max_drawdown_pct": _max_drawdown_pct(equity_curve),
"equity_curve": equity_curve,
},
}