Coverage for src / monte_neo / core / mlx_engine.py: 96%

185 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-28 16:27 +0200

1 

2from __future__ import annotations 

3 

4import logging 

5import os 

6import subprocess 

7import time 

8from typing import TYPE_CHECKING, Any 

9 

10import mlx.core as mx 

11import numpy as np 

12import pandas as pd 

13 

14from monte_neo.core.acceleration.engine import GpuAccelerationEngine 

15from monte_neo.core.gpu_lazy import backtest_lazy_scenarios as run_lazy_backtest 

16from monte_neo.core.gpu_scenarios import normalize_signal_array, run_scenarios_backtest 

17from monte_neo.metrics.calculator import MetricsCalculator 

18from monte_neo.monte_carlo.workers import run_indicator_batch 

19from monte_neo.utils.cache import load_cache, save_cache 

20from monte_neo.utils.parallel import ParallelExecutor 

21 

22logger = logging.getLogger(__name__) 

23 

24if TYPE_CHECKING: 

25 from monte_neo.indicators.base import BaseIndicator 

26 

27# Try to import the Metal bridge extension 

28try: 

29 from monte_neo.core.acceleration.cpp_metal.metal_engine import Candle, Driver, MetalBacktestBridge 

30 METAL_EXTENSION_AVAILABLE = True 

31except ImportError: 

32 # Try to auto-compile if extension is missing 

33 logger.info("🛠️ Metal extension not found. Attempting auto-compilation...") 

34 try: 

35 script_path = os.path.join(os.path.dirname(__file__), "acceleration/cpp_metal/compile.sh") 

36 if os.path.exists(script_path): 

37 result = subprocess.run(["bash", script_path], capture_output=True, text=True) 

38 if result.returncode == 0: 

39 from monte_neo.core.acceleration.cpp_metal.metal_engine import Candle, Driver, MetalBacktestBridge 

40 METAL_EXTENSION_AVAILABLE = True 

41 logger.info("✅ Metal extension compiled and loaded successfully.") 

42 else: 

43 logger.warning(f"❌ Auto-compilation failed: {result.stderr}") 

44 METAL_EXTENSION_AVAILABLE = False 

45 else: 

46 METAL_EXTENSION_AVAILABLE = False 

47 except Exception as e: 

48 logger.warning(f"⚠️ Failed to auto-compile Metal extension: {e}") 

49 METAL_EXTENSION_AVAILABLE = False 

50 

51 

52class MLXBacktestEngine: 

53 """GPU-accelerated backtesting engine using MLX.""" 

54 

55 def __init__(self, precision: str = "float32", metal_driver: str = "cpp", initial_capital: float = 100000.0, leverage: float = 1.0) -> None: 

56 self.precision = precision 

57 self.metal_driver = metal_driver 

58 self.initial_capital = initial_capital 

59 self.leverage = leverage 

60 self.pure_gpu_engine = GpuAccelerationEngine( 

61 precision=precision, 

62 metal_driver=metal_driver, 

63 initial_capital=initial_capital, 

64 leverage=leverage 

65 ) 

66 

67 self.native_bridge = None 

68 if METAL_EXTENSION_AVAILABLE: 

69 if metal_driver == "auto": 

70 self.metal_driver = self._select_best_driver() 

71 else: 

72 self.metal_driver = metal_driver 

73 

74 driver_enum = Driver.CPP 

75 if self.metal_driver == "objc": driver_enum = Driver.OBJC 

76 elif self.metal_driver == "swift": driver_enum = Driver.SWIFT 

77 

78 self.native_bridge = MetalBacktestBridge(driver_enum) 

79 if not self.native_bridge.init(): 

80 logger.warning(f"Failed to initialize native Metal bridge with driver {self.metal_driver}") 

81 self.native_bridge = None 

82 else: 

83 if metal_driver == "auto": 

84 self.metal_driver = "cpp" 

85 else: 

86 self.metal_driver = metal_driver 

87 

88 def _select_best_driver(self) -> str: 

89 """Run a micro-benchmark to select the best Metal driver.""" 

90 if not METAL_EXTENSION_AVAILABLE: 

91 return "cpp" 

92 

93 # Check cache 

94 cached_driver = load_cache("best_metal_driver.json") 

95 if cached_driver: 

96 logger.info(f"🚀 Using cached best Metal driver: {cached_driver}") 

97 return cached_driver 

98 

99 logger.info("🔍 Running micro-benchmark to select best Metal driver...") 

100 candles = [Candle(100.0, 101.0, 99.0, 100.0, 1000.0) for _ in range(1000)] 

101 params = [14.0, 14.0, 1.5, 3.0, 2.0] * 10000 

102 n_scenarios = 10000 

103 

104 best_driver = "cpp" 

105 min_time = float('inf') 

106 

107 for d_name, d_enum in [("cpp", Driver.CPP), ("objc", Driver.OBJC), ("swift", Driver.SWIFT)]: 

108 try: 

109 bridge = MetalBacktestBridge(d_enum) 

110 if bridge.init(): 

111 bridge.run_backtest(candles, params, 1000) 

112 start = time.perf_counter() 

113 bridge.run_backtest(candles, params, n_scenarios) 

114 duration = time.perf_counter() - start 

115 if duration < min_time: 

116 min_time = duration 

117 best_driver = d_name 

118 except Exception as e: 

119 logger.debug(f" Driver {d_name} failed benchmark: {e}") 

120 

121 logger.info(f"✅ Selected best Metal driver: {best_driver} ({1.0/min_time*n_scenarios:.0f} scenarios/sec)") 

122 

123 # Save to cache 

124 save_cache("best_metal_driver.json", best_driver) 

125 

126 return best_driver 

127 

128 def run_full_simulation( 

129 self, 

130 data: pd.DataFrame, 

131 indicator: BaseIndicator, 

132 n_scenarios: int, 

133 method: str = "shuffling", 

134 seed: int = 42, 

135 use_sl_tp: bool = False, 

136 sl_pct: float = 0.0, 

137 tp_pct: float = 0.0, 

138 **kwargs: Any, 

139 ) -> tuple[list[dict[str, Any]], dict[str, float]]: 

140 start_total = time.perf_counter() 

141 timing_stats = {} 

142 mlx_strategy = indicator.to_mlx_representation() 

143 

144 if self.native_bridge and method == "shuffling" and hasattr(indicator, "get_metal_params"): 

145 comm = kwargs.get("commission_bps", 5.0) 

146 slip = kwargs.get("slippage_bps", 5.0) 

147 metal_params = indicator.get_metal_params(commission_bps=comm, slippage_bps=slip) 

148 if metal_params is not None: 

149 try: 

150 t_prep_start = time.perf_counter() 

151 candles = [ 

152 Candle(float(o), float(h), float(l), float(c), float(v)) 

153 for o, h, l, c, v in zip(data['open'], data['high'], data['low'], data['close'], data['volume']) 

154 ] 

155 full_params = metal_params * n_scenarios 

156 timing_stats["data_prep"] = time.perf_counter() - t_prep_start 

157 

158 t_kernel_start = time.perf_counter() 

159 results = self.native_bridge.run_backtest(candles, full_params, n_scenarios) 

160 timing_stats["kernel_execution"] = time.perf_counter() - t_kernel_start 

161 

162 t_format_start = time.perf_counter() 

163 formatted_results = [] 

164 for res in results: 

165 formatted_results.append({ 

166 "metrics": { 

167 "total_return": res.total_return, 

168 "trade_count": res.trade_count, 

169 "profit_factor": res.profit_factor, 

170 "win_rate": res.win_rate, 

171 "max_drawdown": res.max_drawdown, 

172 "sharpe_ratio": res.sharpe_ratio 

173 } 

174 }) 

175 timing_stats["result_formatting"] = time.perf_counter() - t_format_start 

176 timing_stats["total"] = time.perf_counter() - start_total 

177 return formatted_results, timing_stats 

178 except Exception as e: 

179 logger.warning(f"Native Metal bridge execution failed, falling back to MLX: {e}") 

180 

181 if use_sl_tp: 

182 from monte_neo.core.acceleration.tensor_ops import to_tensor 

183 tensors = to_tensor(data) 

184 close = tensors["close"] 

185 from monte_neo.core.acceleration.tensor_ops import TensorOps 

186 if method == "shuffling": 

187 scenarios = TensorOps.generate_shuffle_scenarios(close, n_scenarios, seed=seed) 

188 else: 

189 scenarios = TensorOps.generate_noise_scenarios(close, n_scenarios, std_dev=kwargs.get('std_dev', 0.01), seed=seed) 

190 

191 if mlx_strategy is None: 

192 raise ValueError("Could not create MLX strategy for indicator") 

193 

194 signals = mlx_strategy.generate_signals(scenarios) 

195 scenarios_np = np.array(scenarios).astype(np.float64) 

196 signals_np = np.array(signals).astype(np.int32) 

197 

198 batch_metrics = MetricsCalculator.calculate_batch_multi_price_fast( 

199 scenarios_np, scenarios_np, scenarios_np, signals_np, use_sl_tp, sl_pct, tp_pct 

200 ) 

201 

202 results = [] 

203 for i in range(n_scenarios): 

204 results.append({ 

205 "total_return": float(batch_metrics[i, 0]), 

206 "max_drawdown": float(batch_metrics[i, 1]), 

207 "profit_factor": float(batch_metrics[i, 2]), 

208 "metrics": { 

209 "total_return": float(batch_metrics[i, 0]), 

210 "max_drawdown": float(batch_metrics[i, 1]), 

211 "profit_factor": float(batch_metrics[i, 2]), 

212 "trade_count": int(batch_metrics[i, 3]), 

213 } 

214 }) 

215 timing_stats["total"] = time.perf_counter() - start_total 

216 return results, timing_stats 

217 

218 results = self.pure_gpu_engine.run_simulation( 

219 data=data, mlx_strategy=mlx_strategy, n_scenarios=n_scenarios, method=method, seed=seed 

220 ) 

221 timing_stats["total"] = time.perf_counter() - start_total 

222 return results, timing_stats 

223 

224 def backtest_batch( 

225 self, 

226 data: pd.DataFrame, 

227 indicators: list[BaseIndicator], 

228 executor: ParallelExecutor | None = None, 

229 use_shared_data: bool = False, 

230 force_parallel: bool = False, 

231 parallel_threshold: int = 1000, 

232 dynamic_parallel_threshold: int = 10, 

233 use_sl_tp: bool = False, 

234 sl_pct: float = 0.0, 

235 tp_pct: float = 0.0, 

236 ) -> list[dict[str, Any]]: 

237 close_prices = mx.array(data["close"].to_numpy().astype(np.float32)) 

238 

239 use_parallel = False 

240 if executor and executor.use_processes: 

241 has_dynamic = any(hasattr(ind, "_compile_if_needed") for ind in indicators) 

242 use_parallel = force_parallel or len(indicators) >= parallel_threshold or (has_dynamic and len(indicators) >= dynamic_parallel_threshold) 

243 

244 if use_parallel and executor: 

245 n_workers = executor.n_workers 

246 chunk_size = max(1, len(indicators) // n_workers) 

247 chunks = [indicators[i : i + chunk_size] for i in range(0, len(indicators), chunk_size)] 

248 task_data = None if (use_shared_data and getattr(executor, "initializer", None) is not None) else data 

249 tasks = [(chunk, task_data) for chunk in chunks] 

250 batch_results = executor.map(run_indicator_batch, tasks) 

251 raw_signals = [] 

252 for batch in batch_results: raw_signals.extend(batch) 

253 else: 

254 raw_signals = [ind.generate_signals_fast(data) for ind in indicators] 

255 

256 if use_sl_tp: 

257 close_prices_np = data["close"].to_numpy().astype(np.float64) 

258 high_prices_np = data["high"].to_numpy().astype(np.float64) 

259 low_prices_np = data["low"].to_numpy().astype(np.float64) 

260 signal_list = [normalize_signal_array(sigs, len(data)).astype(np.int32) for sigs in raw_signals] 

261 signal_matrix = np.stack(signal_list) 

262 batch_metrics = MetricsCalculator.calculate_batch_fast( 

263 close_prices_np, high_prices_np, low_prices_np, signal_matrix, use_sl_tp, sl_pct, tp_pct 

264 ) 

265 results = [] 

266 for i in range(len(indicators)): 

267 results.append({ 

268 "total_return": float(batch_metrics[i, 0]), 

269 "max_drawdown": float(batch_metrics[i, 1]), 

270 "profit_factor": float(batch_metrics[i, 2]), 

271 "metrics": { 

272 "total_return": float(batch_metrics[i, 0]), 

273 "max_drawdown": float(batch_metrics[i, 1]), 

274 "profit_factor": float(batch_metrics[i, 2]), 

275 "trade_count": int(batch_metrics[i, 3]), 

276 } 

277 }) 

278 return results 

279 

280 signal_list = [normalize_signal_array(sigs, len(data)) for sigs in raw_signals] 

281 signal_matrix_mx = mx.array(np.stack(signal_list).astype(np.int32)) 

282 returns_pct = (close_prices[1:] / close_prices[:-1]) - 1 

283 strat_returns = signal_matrix_mx[:, :-1] * returns_pct 

284 equity_curves = mx.exp(mx.cumsum(mx.log1p(mx.clip(strat_returns, -0.999, 10.0)), axis=1)) 

285 final_returns = np.array(equity_curves[:, -1]) 

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

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

288 sig_diff = mx.abs(signal_matrix_mx[:, 1:] - signal_matrix_mx[:, :-1]) 

289 trade_counts = np.array(mx.sum(sig_diff > 0, axis=1) / 2) 

290 wins = mx.where(strat_returns > 0, strat_returns, 0) 

291 losses = mx.where(strat_returns < 0, strat_returns, 0) 

292 gross_profit = mx.sum(wins, axis=1) 

293 gross_loss = mx.abs(mx.sum(losses, axis=1)) 

294 pf_np = np.array(mx.where(gross_loss > 0, gross_profit / gross_loss, 100.0)) 

295 

296 results = [] 

297 for i in range(len(indicators)): 

298 results.append({ 

299 "total_return": float(final_returns[i]) - 1.0, 

300 "max_drawdown": float(max_dds[i]), 

301 "profit_factor": float(pf_np[i]), 

302 "metrics": { 

303 "total_return": float(final_returns[i]) - 1.0, 

304 "max_drawdown": float(max_dds[i]), 

305 "profit_factor": float(pf_np[i]), 

306 "trade_count": int(trade_counts[i]), 

307 }, 

308 }) 

309 return results 

310 

311 def backtest_scenarios(self, indicator: BaseIndicator, scenarios: list[pd.DataFrame], executor: ParallelExecutor | None = None, 

312 use_sl_tp: bool = False, sl_pct: float = 0.0, tp_pct: float = 0.0) -> list[dict[str, Any]]: 

313 return run_scenarios_backtest(indicator, scenarios, executor, use_sl_tp=use_sl_tp, sl_pct=sl_pct, tp_pct=tp_pct) 

314 

315 def backtest_lazy_scenarios(self, indicator: BaseIndicator, n_scenarios: int, executor: ParallelExecutor, 

316 block_size: int | None = None, base_seed: int = 42, use_sl_tp: bool = False, 

317 sl_pct: float = 0.0, tp_pct: float = 0.0) -> list[dict[str, Any]]: 

318 return run_lazy_backtest(indicator=indicator, n_scenarios=n_scenarios, executor=executor, block_size=block_size, 

319 base_seed=base_seed, use_sl_tp=use_sl_tp, sl_pct=sl_pct, tp_pct=tp_pct)