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
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
2from __future__ import annotations
4import logging
5import os
6import subprocess
7import time
8from typing import TYPE_CHECKING, Any
10import mlx.core as mx
11import numpy as np
12import pandas as pd
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
22logger = logging.getLogger(__name__)
24if TYPE_CHECKING:
25 from monte_neo.indicators.base import BaseIndicator
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
52class MLXBacktestEngine:
53 """GPU-accelerated backtesting engine using MLX."""
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 )
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
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
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
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"
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
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
104 best_driver = "cpp"
105 min_time = float('inf')
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}")
121 logger.info(f"✅ Selected best Metal driver: {best_driver} ({1.0/min_time*n_scenarios:.0f} scenarios/sec)")
123 # Save to cache
124 save_cache("best_metal_driver.json", best_driver)
126 return best_driver
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()
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
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
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}")
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)
191 if mlx_strategy is None:
192 raise ValueError("Could not create MLX strategy for indicator")
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)
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 )
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
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
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))
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)
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]
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
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))
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
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)
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)