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

1"""Lazy GPU backtesting helpers.""" 

2 

3from __future__ import annotations 

4 

5from typing import TYPE_CHECKING, Any 

6 

7import mlx.core as mx 

8import numpy as np 

9 

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 

13 

14if TYPE_CHECKING: 

15 from monte_neo.indicators.base import BaseIndicator 

16 

17 

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). 

29 

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. 

39 

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] 

45 

46 raw_results = executor.map(_generate_lazy_scenario_wrapper, tasks) 

47 

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 

51 

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 [] 

56 

57 # 2. Prepare Data and Signal Matrix 

58 max_len = max(len(res[2]) for res in valid_data) 

59 

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) 

64 

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] 

71 

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) 

77 

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 

83 

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 ) 

94 

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]) 

101 

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 

115 

116 signal_list = [] 

117 returns_list = [] 

118 

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)) 

127 

128 if not signal_list: 

129 return [] 

130 

131 max_len = max(len(s) for s in signal_list) 

132 padded_signals = [] 

133 padded_returns = [] 

134 

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) 

143 

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)) 

147 

148 strat_returns = signal_matrix_mx * returns_matrix 

149 

150 equity_curves = mx.exp( 

151 mx.cumsum(mx.log1p(mx.clip(strat_returns, -0.9, 10.0)), axis=1) 

152 ) 

153 

154 final_rets = np.array(equity_curves[:, -1]) 

155 

156 running_max = mx.cummax(equity_curves, axis=1) 

157 max_dds = np.array(mx.max((running_max - equity_curves) / running_max, axis=1)) 

158 

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)) 

164 

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