Coverage for src / monte_neo / core / gpu_lazy.py: 94%
78 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
1"""Lazy GPU backtesting helpers."""
3from __future__ import annotations
5from typing import TYPE_CHECKING, Any
7import mlx.core as mx
8import numpy as np
10from monte_neo.metrics.calculator import MetricsCalculator
11from monte_neo.monte_carlo.workers import _generate_lazy_scenario_wrapper
12from monte_neo.utils.parallel import ParallelExecutor
14if TYPE_CHECKING:
15 from monte_neo.indicators.base import BaseIndicator
18def backtest_lazy_scenarios(
19 indicator: BaseIndicator,
20 n_scenarios: int,
21 executor: ParallelExecutor,
22 block_size: int | None = None,
23 base_seed: int = 42,
24 use_sl_tp: bool = False,
25 sl_pct: float = 0.0,
26 tp_pct: float = 0.0,
27) -> list[dict[str, Any]]:
28 """Run backtest with lazy scenario generation (Block Bootstrap).
30 Args:
31 indicator: Indicator to test.
32 n_scenarios: Number of scenarios.
33 executor: Shared parallel executor.
34 block_size: Optional block size for bootstrap.
35 base_seed: Base seed for reproducible scenarios.
36 use_sl_tp: Whether to apply Stop Loss and Take Profit.
37 sl_pct: Stop Loss percentage.
38 tp_pct: Take Profit percentage.
40 Returns:
41 List of result dictionaries.
42 """
43 seeds = [base_seed + i for i in range(n_scenarios)]
44 tasks = [(indicator, seed, block_size) for seed in seeds]
46 raw_results = executor.map(_generate_lazy_scenario_wrapper, tasks)
48 if use_sl_tp:
49 # Optimization: Use parallelized batch Numba for SL/TP on multiple scenarios
50 # This is much faster than the previous Python loop
52 # 1. Extract valid results (sigs, rets, ohlc)
53 valid_data = [res for res in raw_results if res is not None]
54 if not valid_data:
55 return []
57 # 2. Prepare Data and Signal Matrix
58 max_len = max(len(res[2]) for res in valid_data)
60 close_matrix = np.zeros((len(valid_data), max_len), dtype=np.float64)
61 high_matrix = np.zeros((len(valid_data), max_len), dtype=np.float64)
62 low_matrix = np.zeros((len(valid_data), max_len), dtype=np.float64)
63 signal_matrix = np.zeros((len(valid_data), max_len), dtype=np.int32)
65 for i, (sigs, _, ohlc) in enumerate(valid_data):
66 l = len(ohlc)
67 # ohlc is [open, high, low, close]
68 high_matrix[i, :l] = ohlc[:, 1]
69 low_matrix[i, :l] = ohlc[:, 2]
70 close_matrix[i, :l] = ohlc[:, 3]
72 # Normalize signals
73 if hasattr(sigs, "to_numpy"):
74 s_arr = sigs["signal"].to_numpy() if "signal" in sigs.columns else sigs.to_numpy().reshape(-1)
75 else:
76 s_arr = np.asarray(sigs).reshape(-1)
78 s_arr = s_arr.astype(np.int32, copy=False)
79 if s_arr.size >= l:
80 signal_matrix[i, :l] = s_arr[:l]
81 else:
82 signal_matrix[i, :s_arr.size] = s_arr
84 # 3. Run Batch Calculation
85 batch_metrics = MetricsCalculator.calculate_batch_multi_price_fast(
86 close_matrix,
87 high_matrix,
88 low_matrix,
89 signal_matrix,
90 use_sl_tp,
91 sl_pct,
92 tp_pct
93 )
95 results = []
96 for i in range(len(valid_data)):
97 total_return = float(batch_metrics[i, 0])
98 max_dd = float(batch_metrics[i, 1])
99 pf = float(batch_metrics[i, 2])
100 trade_count = int(batch_metrics[i, 3])
102 results.append({
103 "total_return": total_return,
104 "max_drawdown": max_dd,
105 "profit_factor": pf,
106 "passed": bool(total_return > 0.0 and max_dd < 0.2),
107 "metrics": {
108 "total_return": total_return,
109 "max_drawdown": max_dd,
110 "profit_factor": pf,
111 "trade_count": trade_count,
112 }
113 })
114 return results
116 signal_list = []
117 returns_list = []
119 for res in raw_results:
120 if res is None:
121 continue
122 sigs, rets, _ = res
123 sig_vals = sigs.reshape(-1).astype(np.float32)
124 sig_vals = sig_vals[:-1]
125 signal_list.append(sig_vals)
126 returns_list.append(rets.astype(np.float32))
128 if not signal_list:
129 return []
131 max_len = max(len(s) for s in signal_list)
132 padded_signals = []
133 padded_returns = []
135 for s, r in zip(signal_list, returns_list):
136 pad_width = max_len - len(s)
137 if pad_width > 0:
138 padded_signals.append(np.pad(s, (0, pad_width), "constant", constant_values=0))
139 padded_returns.append(np.pad(r, (0, pad_width), "constant", constant_values=0))
140 else:
141 padded_signals.append(s)
142 padded_returns.append(r)
144 signal_matrix_np = np.stack(padded_signals)
145 signal_matrix_mx: Any = mx.array(signal_matrix_np.astype(np.int32))
146 returns_matrix = mx.array(np.stack(padded_returns))
148 strat_returns = signal_matrix_mx * returns_matrix
150 equity_curves = mx.exp(
151 mx.cumsum(mx.log1p(mx.clip(strat_returns, -0.9, 10.0)), axis=1)
152 )
154 final_rets = np.array(equity_curves[:, -1])
156 running_max = mx.cummax(equity_curves, axis=1)
157 max_dds = np.array(mx.max((running_max - equity_curves) / running_max, axis=1))
159 wins = mx.where(strat_returns > 0, strat_returns, 0)
160 losses = mx.where(strat_returns < 0, strat_returns, 0)
161 gross_profit = mx.sum(wins, axis=1)
162 gross_loss = mx.abs(mx.sum(losses, axis=1))
163 profit_factor = np.array(mx.where(gross_loss > 0, gross_profit / gross_loss, 100.0))
165 results = []
166 for i in range(len(signal_list)):
167 results.append(
168 {
169 "total_return": float(final_rets[i]) - 1.0,
170 "max_drawdown": float(max_dds[i]),
171 "profit_factor": float(profit_factor[i]),
172 "passed": bool(final_rets[i] > 1.0 and max_dds[i] < 0.2),
173 "metrics": {
174 "total_return": float(final_rets[i]) - 1.0,
175 "max_drawdown": float(max_dds[i]),
176 "profit_factor": float(profit_factor[i]),
177 },
178 }
179 )
180 return results