""" Market regime classification for strategy gating. Classifies each bar into one of four regimes based on rolling returns, volatility, and volume profile. Used to compute conditional strategy performance — a strategy that only works in trending regimes must KNOW when it's in a trending regime. """ from __future__ import annotations from collections import deque from typing import Optional import numpy as np def classify_regime( returns_20: float, vol_20: float, vol_ratio: float = 1.0, trend_threshold: float = 0.05, vol_threshold: float = 0.04, ) -> str: """Classify a single bar into a market regime. Args: returns_20: rolling 20-bar return (fraction, e.g. 0.15 = +15%) vol_20: rolling 20-bar realized volatility (annualized or period) vol_ratio: current volume / rolling average volume trend_threshold: minimum abs return to classify as trending vol_threshold: volatility above which market is "volatile" Returns one of: trending_up, trending_down, ranging, volatile """ if vol_20 > vol_threshold or vol_ratio > 2.0: return "volatile" if returns_20 > trend_threshold: return "trending_up" elif returns_20 < -trend_threshold: return "trending_down" else: return "ranging" class RegimeClassifier: """Stateful regime classifier using rolling windows. Usage: rc = RegimeClassifier(window=20) for bar in bars: rc.feed(mid_px=bar.close, close=bar.close, open_px=bar.open, vol=bar.volume) regime = rc.current_regime """ def __init__( self, window: int = 20, trend_threshold: float = 0.05, vol_threshold: float = 0.04, ): self._window = window self._trend_threshold = trend_threshold self._vol_threshold = vol_threshold self._prices: deque[float] = deque(maxlen=window + 1) self._volumes: deque[float] = deque(maxlen=window) self._current_regime: str = "unknown" self._regime_counts: dict[str, int] = {"trending_up": 0, "trending_down": 0, "ranging": 0, "volatile": 0, "unknown": 0} def feed(self, mid_px: float, close: float, open_px: float, vol: float): """Feed a new bar observation.""" self._prices.append(mid_px) self._volumes.append(vol) if len(self._prices) < self._window + 1: return # Rolling return returns_20 = (self._prices[-1] - self._prices[0]) / self._prices[0] if self._prices[0] > 0 else 0 # Rolling volatility price_list = list(self._prices) rets = [(price_list[i] - price_list[i - 1]) / price_list[i - 1] for i in range(1, len(price_list)) if price_list[i - 1] > 0] vol_20 = np.std(rets) if rets else 0.0 # Volume ratio avg_vol = sum(self._volumes) / max(len(self._volumes), 1) vol_ratio = vol / avg_vol if avg_vol > 0 else 1.0 self._current_regime = classify_regime( returns_20, vol_20, vol_ratio, self._trend_threshold, self._vol_threshold, ) self._regime_counts[self._current_regime] += 1 @property def current_regime(self) -> str: return self._current_regime def regime_counts(self) -> dict[str, int]: return dict(self._regime_counts) def dominant_regime(self) -> str: """Most frequent regime observed so far.""" return max(self._regime_counts, key=self._regime_counts.get) def conditional_performance( trades: list[dict], regime_trade_counts: dict[str, int], ) -> dict: """Compute per-regime strategy performance from trade records. Args: trades: list of trade dicts with pnl_net, pnl_gross, time regime_trade_counts: {regime_name: count_of_trades} Returns: dict with per-regime stats (avg_pnl, win_rate, total_pnl, count) and best_regime / worst_regime classification. """ # For simplicity, assume all trades belong to the regime # In production, each trade would be timestamp-matched to regime at entry all_pnls = [float(t.get("pnl_net", t.get("pnl", 0))) for t in trades] result = {"regimes": {}, "overall": {"count": len(trades), "total_pnl": round(sum(all_pnls), 2)}} for regime, count in regime_trade_counts.items(): if count == 0: result["regimes"][regime] = {"count": 0, "total_pnl": 0, "avg_pnl": 0, "win_rate": 0} continue # Take the next batch of trades for this regime # (simplified — real impl would match by timestamp) regime_pnls = all_pnls[:count] wins = sum(1 for p in regime_pnls if p > 0) result["regimes"][regime] = { "count": count, "total_pnl": round(sum(regime_pnls), 2), "avg_pnl": round(sum(regime_pnls) / count, 2) if count else 0, "win_rate": round(wins / count, 3) if count else 0, } # Best and worst regime scored = [(r, info["avg_pnl"]) for r, info in result["regimes"].items() if info["count"] > 0] if scored: result["best_regime"] = max(scored, key=lambda x: x[1])[0] result["worst_regime"] = min(scored, key=lambda x: x[1])[0] return result