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

1"""Unified metrics calculator. 

2 

3Calculates all trading metrics from signals and data. 

4""" 

5 

6from __future__ import annotations 

7 

8import numpy as np 

9import pandas as pd 

10 

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 

20 

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 

27 

28logger = get_logger(__name__) 

29 

30 

31class MetricsCalculator: 

32 """Calculate all trading metrics.""" 

33 

34 def __init__(self, risk_free_rate: float = 0.0, initial_capital: float = 100000.0, leverage: float = 1.0) -> None: 

35 """Initialize metrics calculator. 

36 

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 

45 

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

52 

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. 

65 

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

75 

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) 

81 

82 if not trades: 

83 return metric_utils.get_empty_metrics() 

84 

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] 

88 

89 metrics = {} 

90 

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

94 

95 def needs(name: str) -> bool: 

96 return need_all or name in reqs 

97 

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) 

127 

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 } 

137 

138 if need_all or not reqs.isdisjoint(equity_metrics): 

139 equity = self._calculate_equity(trades) 

140 max_dd_val = 0.0 

141 

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 

150 

151 if needs("avg_drawdown"): 

152 metrics["avg_drawdown"] = self.drawdown.calculate_avg(equity) 

153 

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 ) 

159 

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 ) 

165 

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

169 

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) 

174 

175 return metrics 

176 

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] 

199 

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) 

207 

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 ] 

226 

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 ) 

239 

240 return [TradeResult(*t) for t in raw_trades] 

241 

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 ) 

258 

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 ) 

283 

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

288 

289 pnl_pcts = np.array([t.pnl_pct for t in trades]) 

290 # Apply leverage 

291 effective_pnls = pnl_pcts * self.leverage 

292 

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) 

296 

297 return equity