Coverage for src / monte_neo / metrics / calculator.py: 98%
122 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"""Unified metrics calculator.
3Calculates all trading metrics from signals and data.
4"""
6from __future__ import annotations
8import numpy as np
9import pandas as pd
11from monte_neo.core import native_metrics # type: ignore
12from monte_neo.metrics import numba_funcs
13from monte_neo.metrics import utils as metric_utils
14from monte_neo.metrics.drawdown import DrawdownMetric
15from monte_neo.metrics.profit_factor import ProfitFactorMetric
16from monte_neo.metrics.sharpe import SharpeRatioMetric, SortinoRatioMetric
17from monte_neo.metrics.types import TradeResult
18from monte_neo.metrics.winrate import WinrateMetric
19from monte_neo.utils.logger import get_logger
21try:
22 # Check if native module is available and working
23 native_metrics.extract_trades
24 HAS_NATIVE = True
25except (ImportError, AttributeError):
26 HAS_NATIVE = False
28logger = get_logger(__name__)
31class MetricsCalculator:
32 """Calculate all trading metrics."""
34 def __init__(self, risk_free_rate: float = 0.0, initial_capital: float = 100000.0, leverage: float = 1.0) -> None:
35 """Initialize metrics calculator.
37 Args:
38 risk_free_rate: Annual risk-free rate for Sharpe calculation.
39 initial_capital: Initial account balance.
40 leverage: Trading leverage (default 1.0 = no leverage).
41 """
42 self.risk_free_rate = risk_free_rate
43 self.initial_capital = initial_capital
44 self.leverage = leverage
46 # Initialize individual metric calculators
47 self.profit_factor = ProfitFactorMetric()
48 self.sharpe = SharpeRatioMetric(risk_free_rate)
49 self.sortino = SortinoRatioMetric(risk_free_rate)
50 self.drawdown = DrawdownMetric()
51 self.winrate = WinrateMetric()
53 def calculate_all(
54 self,
55 data: pd.DataFrame | np.ndarray,
56 signals: pd.DataFrame | np.ndarray,
57 required_metrics: list[str] | None = None,
58 use_sl_tp: bool = False,
59 sl_pct: float = 0.0,
60 tp_pct: float = 0.0,
61 commission_pct: float = 0.0,
62 slippage_pct: float = 0.0,
63 ) -> dict[str, float]:
64 """Calculate metrics.
66 Args:
67 data: OHLCV DataFrame or numpy array (OHLCV).
68 signals: DataFrame with 'signal' column or numpy array of signals.
69 required_metrics: Optional list of metrics to calculate. If None, calculate all.
70 use_sl_tp: Whether to apply Stop Loss and Take Profit.
71 sl_pct: Stop Loss percentage (e.g., 1.0 for 1%).
72 tp_pct: Take Profit percentage (e.g., 2.0 for 2%).
73 commission_pct: Commission percentage per trade (e.g., 0.05 for 0.05%).
74 slippage_pct: Slippage percentage per trade (e.g., 0.05 for 0.05%).
76 Returns:
77 Dictionary of metrics.
78 """
79 # Extract trades from signals
80 trades = self._extract_trades(data, signals, use_sl_tp, sl_pct, tp_pct, commission_pct, slippage_pct)
82 if not trades:
83 return metric_utils.get_empty_metrics()
85 # Basic PnLs are needed for almost everything
86 pnls = [t.pnl for t in trades]
87 pnl_pcts = [t.pnl_pct for t in trades]
89 metrics = {}
91 # If required_metrics is provided, check what we need
92 need_all = required_metrics is None
93 reqs = set(required_metrics) if required_metrics else set()
95 def needs(name: str) -> bool:
96 return need_all or name in reqs
98 # Profit metrics
99 if needs("profit_factor"):
100 metrics["profit_factor"] = self.profit_factor.calculate(pnls)
101 if needs("total_return"):
102 metrics["total_return"] = float(np.sum(pnl_pcts))
103 if needs("total_profit_abs"):
104 equity = self._calculate_equity(trades)
105 metrics["total_profit_abs"] = float(equity[-1] - self.initial_capital)
106 if needs("final_balance"):
107 equity = self._calculate_equity(trades)
108 metrics["final_balance"] = float(equity[-1])
109 if needs("avg_return"):
110 metrics["avg_return"] = float(np.mean(pnl_pcts)) if pnl_pcts else 0.0
111 if needs("winrate"):
112 metrics["winrate"] = self.winrate.calculate(pnls)
113 if needs("expectancy"):
114 metrics["expectancy"] = self.winrate.expectancy(pnls)
115 if needs("avg_win"):
116 metrics["avg_win"] = self.winrate.avg_win(pnls)
117 if needs("avg_loss"):
118 metrics["avg_loss"] = self.winrate.avg_loss(pnls)
119 if needs("win_loss_ratio"):
120 metrics["win_loss_ratio"] = self.winrate.win_loss_ratio(pnls)
121 if needs("trade_count"):
122 metrics["trade_count"] = len(trades)
123 if needs("consecutive_wins"):
124 metrics["consecutive_wins"] = metric_utils.max_consecutive(pnls, True)
125 if needs("consecutive_losses"):
126 metrics["consecutive_losses"] = metric_utils.max_consecutive(pnls, False)
128 # Complex metrics requiring Equity Curve
129 equity_metrics = {
130 "sharpe_ratio",
131 "sortino_ratio",
132 "max_drawdown",
133 "avg_drawdown",
134 "recovery_factor",
135 "calmar_ratio",
136 }
138 if need_all or not reqs.isdisjoint(equity_metrics):
139 equity = self._calculate_equity(trades)
140 max_dd_val = 0.0
142 if (
143 needs("max_drawdown")
144 or needs("recovery_factor")
145 or needs("calmar_ratio")
146 ):
147 max_dd_val = self.drawdown.calculate_max(equity)
148 if needs("max_drawdown"):
149 metrics["max_drawdown"] = max_dd_val
151 if needs("avg_drawdown"):
152 metrics["avg_drawdown"] = self.drawdown.calculate_avg(equity)
154 if needs("recovery_factor"):
155 total_ret = float(np.sum(pnl_pcts))
156 metrics["recovery_factor"] = metric_utils.calculate_recovery_factor(
157 total_ret, max_dd_val
158 )
160 if needs("calmar_ratio"):
161 avg_ret = float(np.mean(pnl_pcts)) if pnl_pcts else 0.0
162 metrics["calmar_ratio"] = metric_utils.calculate_calmar_ratio(
163 avg_ret, max_dd_val
164 )
166 # Returns based metrics
167 if needs("sharpe_ratio") or needs("sortino_ratio"):
168 returns = np.diff(equity) / equity[:-1] if len(equity) > 1 else []
170 if needs("sharpe_ratio"):
171 metrics["sharpe_ratio"] = self.sharpe.calculate(returns)
172 if needs("sortino_ratio"):
173 metrics["sortino_ratio"] = self.sortino.calculate(returns)
175 return metrics
177 def _extract_trades(
178 self,
179 data: pd.DataFrame | np.ndarray,
180 signals: pd.DataFrame | np.ndarray,
181 use_sl_tp: bool = False,
182 sl_pct: float = 0.0,
183 tp_pct: float = 0.0,
184 commission_pct: float = 0.0,
185 slippage_pct: float = 0.0,
186 ) -> list[TradeResult]:
187 """Extract trades from signals. Use C++ if available."""
188 # Convert data to numpy arrays if it's a DataFrame
189 if isinstance(data, pd.DataFrame):
190 close_prices = data["close"].to_numpy()
191 high_prices = data["high"].to_numpy()
192 low_prices = data["low"].to_numpy()
193 else:
194 # Assume data is a numpy array (OHLCV)
195 # col 1=high, 2=low, 3=close
196 high_prices = data[:, 1]
197 low_prices = data[:, 2]
198 close_prices = data[:, 3]
200 # Convert signals to numpy array if it's a DataFrame
201 if isinstance(signals, pd.DataFrame):
202 if "signal" not in signals.columns:
203 return []
204 signal_array = signals["signal"].to_numpy().astype(np.int32)
205 else:
206 signal_array = np.asarray(signals, dtype=np.int32)
208 if HAS_NATIVE and not use_sl_tp:
209 # Use high-performance C++ extension (native doesn't support SL/TP yet)
210 raw_trades = native_metrics.extract_trades(
211 close_prices.tolist(),
212 signal_array.tolist(),
213 )
214 return [
215 TradeResult(
216 entry_idx=t.entry_idx,
217 exit_idx=t.exit_idx,
218 entry_price=t.entry_price,
219 exit_price=t.exit_price,
220 direction=t.direction,
221 pnl=t.pnl,
222 pnl_pct=t.pnl_pct,
223 )
224 for t in raw_trades
225 ]
227 # Fallback to JIT-compiled Python
228 raw_trades = numba_funcs.extract_trades_fast(
229 close_prices,
230 high_prices,
231 low_prices,
232 signal_array,
233 use_sl_tp,
234 sl_pct,
235 tp_pct,
236 commission_pct,
237 slippage_pct,
238 )
240 return [TradeResult(*t) for t in raw_trades]
242 @staticmethod
243 def calculate_batch_fast(
244 prices: np.ndarray,
245 highs: np.ndarray,
246 lows: np.ndarray,
247 signal_matrix: np.ndarray,
248 use_sl_tp: bool,
249 sl_pct: float,
250 tp_pct: float,
251 commission_pct: float = 0.0,
252 slippage_pct: float = 0.0,
253 ) -> np.ndarray:
254 """Calculate basic metrics for a batch of signal sets in parallel."""
255 return numba_funcs.calculate_batch_fast(
256 prices, highs, lows, signal_matrix, use_sl_tp, sl_pct, tp_pct, commission_pct, slippage_pct
257 )
259 @staticmethod
260 def calculate_batch_multi_price_fast(
261 price_matrix: np.ndarray,
262 high_matrix: np.ndarray,
263 low_matrix: np.ndarray,
264 signal_matrix: np.ndarray,
265 use_sl_tp: bool,
266 sl_pct: float,
267 tp_pct: float,
268 commission_pct: float = 0.0,
269 slippage_pct: float = 0.0,
270 ) -> np.ndarray:
271 """Calculate basic metrics for a batch where each row has its own prices."""
272 return numba_funcs.calculate_batch_multi_price_fast(
273 price_matrix,
274 high_matrix,
275 low_matrix,
276 signal_matrix,
277 use_sl_tp,
278 sl_pct,
279 tp_pct,
280 commission_pct,
281 slippage_pct,
282 )
284 def _calculate_equity(self, trades: list[TradeResult]) -> np.ndarray:
285 """Calculate equity curve from trades using vectorized cumprod."""
286 if not trades:
287 return np.array([self.initial_capital])
289 pnl_pcts = np.array([t.pnl_pct for t in trades])
290 # Apply leverage
291 effective_pnls = pnl_pcts * self.leverage
293 # Equity starts at initial_capital, then cumprod of (1 + pnl_pct)
294 equity = np.ones(len(trades) + 1) * self.initial_capital
295 equity[1:] = self.initial_capital * np.cumprod(1 + effective_pnls)
297 return equity