malkhut(T6): training core — CMA-ES trainer, registry, pipeline, selector
CMA-ES trainer (cma_trainer.py): self-play pool, bootstrap CI, ScenarioFactory with behavior-driven scenarios, auto-compile, label query interfaces. Policy registry (registry.py): CANDIDATE → ACTIVE lifecycle. Training pipeline (pipeline.py): bounded continuous learning loop + logger. Strategy selector (selector.py): regime → strategy mapping, performance matrix.
This commit is contained in:
391
MALKHUT/malkhut/training/selector.py
Normal file
391
MALKHUT/malkhut/training/selector.py
Normal file
@@ -0,0 +1,391 @@
|
||||
"""
|
||||
Strategy Selector — maps market conditions to best strategy.
|
||||
|
||||
Architecture:
|
||||
market_fingerprint (ExoF + MARAS + vel_div + eigenscan)
|
||||
↓
|
||||
Regime Classifier (market state → regime tag)
|
||||
↓
|
||||
Performance Matrix (regime × strategy → score)
|
||||
↓
|
||||
Selector (best strategy for current regime)
|
||||
↓
|
||||
FulfilmentPolicyParams (selected strategy's parameters)
|
||||
|
||||
The system maintains a PORTFOLIO of strategies and SELECTS the best one
|
||||
for current conditions. Strategies are adversarially tested across regimes
|
||||
and floated to the top based on regime-specific performance.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Optional, Sequence, Tuple
|
||||
|
||||
from malkhut.state import FulfilmentPolicyParams, MarketWorldState
|
||||
from malkhut.features import DefaultFeatureExtractor, FeatureExtractor
|
||||
from malkhut.training.cma_trainer import PolicySnapshot
|
||||
from malkhut.training.dsl import StrategyTemplate, SensorType, _read_sensor
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Market Regime Classification
|
||||
# ==============================================================================
|
||||
|
||||
class MarketRegime(str, Enum):
|
||||
"""
|
||||
Market regime tags derived from fingerprinting.
|
||||
|
||||
These map to upstream systems:
|
||||
- MARAS 8-regime classification (ExF/Eigen/BTC/EsoF/Micro)
|
||||
- DOLPHINNG7 vel_div + eigenscan
|
||||
- ExoF external factors
|
||||
|
||||
Each regime has characteristic behavior where different strategies excel.
|
||||
"""
|
||||
TRENDING_UP = "trending_up"
|
||||
TRENDING_DOWN = "trending_down"
|
||||
HIGH_VOLATILITY = "high_volatility"
|
||||
LOW_VOLATILITY = "low_volatility"
|
||||
MEAN_REVERTING = "mean_reverting"
|
||||
MOMENTUM = "momentum"
|
||||
CHOPPY = "choppy"
|
||||
LIQUIDITY_HOLE = "liquidity_hole"
|
||||
NORMAL = "normal"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MarketFingerprint:
|
||||
"""
|
||||
Snapshot of market state features used for regime classification.
|
||||
|
||||
Derived from upstream systems:
|
||||
- ExoF: funding_bps, open_interest_change
|
||||
- MARAS: regime_score (0-1), 17-dim composite signature
|
||||
- DOLPHINNG7: vel_div, eigenscan
|
||||
"""
|
||||
ts_ns: int
|
||||
regime: MarketRegime
|
||||
volatility: float
|
||||
spread_bps: float
|
||||
imbalance: float
|
||||
toxicity: float
|
||||
regime_score: float # 0-1, from MARAS
|
||||
funding_bps: float
|
||||
price_momentum: float
|
||||
volume_ratio: float # current vs average
|
||||
depth_ratio: float # bid/ask depth ratio
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Regime Classifier
|
||||
# ==============================================================================
|
||||
|
||||
class RegimeClassifier:
|
||||
"""
|
||||
Classify current market state into a regime tag.
|
||||
|
||||
Uses simple rule-based classification for Phase 1.
|
||||
Phase 2: integrate with MARAS ensemble for production classification.
|
||||
"""
|
||||
|
||||
def __init__(self, feature_extractor: Optional[FeatureExtractor] = None) -> None:
|
||||
self._extractor = feature_extractor or DefaultFeatureExtractor()
|
||||
|
||||
def classify(self, state: MarketWorldState) -> MarketRegime:
|
||||
"""Classify current market state into a regime."""
|
||||
fv = self._extractor.extract(state).values
|
||||
|
||||
volatility = state.trade_path.volatility_bps if state.trade_path else 0.0
|
||||
spread_bps = fv.get("spread_bps", 0.0)
|
||||
imbalance = fv.get("top5_imbalance", 0.0)
|
||||
toxicity = fv.get("orderflow_toxicity", 0.0)
|
||||
regime_score = state.trade_path.dolphin_regime_score if state.trade_path else 0.5
|
||||
|
||||
# Rule-based classification (check specific first)
|
||||
if spread_bps > 10:
|
||||
return MarketRegime.LIQUIDITY_HOLE
|
||||
elif volatility > 25:
|
||||
return MarketRegime.HIGH_VOLATILITY
|
||||
elif volatility < 5 and spread_bps < 2:
|
||||
return MarketRegime.LOW_VOLATILITY
|
||||
elif abs(imbalance) > 0.4:
|
||||
return MarketRegime.MOMENTUM if imbalance > 0 else MarketRegime.TRENDING_DOWN
|
||||
elif regime_score > 0.7:
|
||||
return MarketRegime.TRENDING_UP
|
||||
elif regime_score < 0.3:
|
||||
return MarketRegime.MEAN_REVERTING
|
||||
elif abs(imbalance) < 0.1 and spread_bps > 5:
|
||||
return MarketRegime.CHOPPY
|
||||
else:
|
||||
return MarketRegime.NORMAL
|
||||
|
||||
def fingerprint(self, state: MarketWorldState) -> MarketFingerprint:
|
||||
"""Create a full market fingerprint."""
|
||||
fv = DefaultFeatureExtractor().extract(state).values
|
||||
return MarketFingerprint(
|
||||
ts_ns=state.ts_ns,
|
||||
regime=self.classify(state),
|
||||
volatility=fv.get("volatility_state", 0.0),
|
||||
spread_bps=fv.get("spread_bps", 0.0),
|
||||
imbalance=fv.get("top5_imbalance", 0.0),
|
||||
toxicity=fv.get("orderflow_toxicity", 0.0),
|
||||
regime_score=state.trade_path.dolphin_regime_score if state.trade_path else 0.5,
|
||||
funding_bps=state.funding_bps or 0.0,
|
||||
price_momentum=0.0, # placeholder
|
||||
volume_ratio=1.0, # placeholder
|
||||
depth_ratio=1.0, # placeholder
|
||||
)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Performance Matrix
|
||||
# ==============================================================================
|
||||
|
||||
@dataclass
|
||||
class RegimeStrategyScore:
|
||||
"""Performance score for a strategy in a specific regime."""
|
||||
strategy_id: str
|
||||
regime: MarketRegime
|
||||
score: float
|
||||
episodes: int
|
||||
avg_pnl_bps: float
|
||||
avg_drawdown_bps: float
|
||||
avg_adverse_fill_ratio: float
|
||||
last_updated_ns: int
|
||||
|
||||
|
||||
class PerformanceMatrix:
|
||||
"""
|
||||
Tracks strategy performance across regimes.
|
||||
|
||||
Matrix: (regime, strategy_id) → RegimeStrategyScore
|
||||
|
||||
Used by the selector to choose the best strategy for current conditions.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._scores: Dict[Tuple[MarketRegime, str], RegimeStrategyScore] = {}
|
||||
self._strategy_regime_history: Dict[str, List[MarketRegime]] = defaultdict(list)
|
||||
|
||||
def record(
|
||||
self,
|
||||
strategy_id: str,
|
||||
regime: MarketRegime,
|
||||
score: float,
|
||||
pnl_bps: float = 0.0,
|
||||
drawdown_bps: float = 0.0,
|
||||
adverse_fill_ratio: float = 0.0,
|
||||
) -> None:
|
||||
"""Record a strategy's performance in a regime."""
|
||||
key = (regime, strategy_id)
|
||||
existing = self._scores.get(key)
|
||||
|
||||
if existing:
|
||||
# Exponential moving average
|
||||
alpha = 0.3
|
||||
new_score = alpha * score + (1 - alpha) * existing.score
|
||||
new_episodes = existing.episodes + 1
|
||||
new_pnl = alpha * pnl_bps + (1 - alpha) * existing.avg_pnl_bps
|
||||
new_dd = alpha * drawdown_bps + (1 - alpha) * existing.avg_drawdown_bps
|
||||
new_adverse = alpha * adverse_fill_ratio + (1 - alpha) * existing.avg_adverse_fill_ratio
|
||||
else:
|
||||
new_score = score
|
||||
new_episodes = 1
|
||||
new_pnl = pnl_bps
|
||||
new_dd = drawdown_bps
|
||||
new_adverse = adverse_fill_ratio
|
||||
|
||||
self._scores[key] = RegimeStrategyScore(
|
||||
strategy_id=strategy_id,
|
||||
regime=regime,
|
||||
score=new_score,
|
||||
episodes=new_episodes,
|
||||
avg_pnl_bps=new_pnl,
|
||||
avg_drawdown_bps=new_dd,
|
||||
avg_adverse_fill_ratio=new_adverse,
|
||||
last_updated_ns=time.time_ns(),
|
||||
)
|
||||
self._strategy_regime_history[strategy_id].append(regime)
|
||||
|
||||
def get_best(
|
||||
self,
|
||||
regime: MarketRegime,
|
||||
exclude: Optional[set[str]] = None,
|
||||
) -> Optional[str]:
|
||||
"""Get the best strategy for a given regime."""
|
||||
exclude = exclude or set()
|
||||
candidates = [
|
||||
(key[1], score.score)
|
||||
for key, score in self._scores.items()
|
||||
if key[0] == regime and key[1] not in exclude
|
||||
]
|
||||
if not candidates:
|
||||
return None
|
||||
return max(candidates, key=lambda x: x[1])[0]
|
||||
|
||||
def get_scores_for_regime(self, regime: MarketRegime) -> List[RegimeStrategyScore]:
|
||||
"""Get all strategy scores for a regime, sorted by score."""
|
||||
scores = [s for s in self._scores.values() if s.regime == regime]
|
||||
return sorted(scores, key=lambda s: s.score, reverse=True)
|
||||
|
||||
def get_regimes_for_strategy(self, strategy_id: str) -> List[MarketRegime]:
|
||||
"""Get all regimes where a strategy has been tested."""
|
||||
return list(set(key[0] for key in self._scores if key[1] == strategy_id))
|
||||
|
||||
def get_coverage(self) -> Dict[str, int]:
|
||||
"""Get regime coverage per strategy."""
|
||||
coverage: Dict[str, int] = defaultdict(int)
|
||||
for key in self._scores:
|
||||
coverage[key[1]] += 1
|
||||
return dict(coverage)
|
||||
|
||||
@property
|
||||
def total_entries(self) -> int:
|
||||
return len(self._scores)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Strategy Selector
|
||||
# ==============================================================================
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SelectionResult:
|
||||
"""Result of strategy selection."""
|
||||
strategy_id: str
|
||||
regime: MarketRegime
|
||||
score: float
|
||||
confidence: float # 0-1, how confident we are in this selection
|
||||
alternatives: Tuple[str, ...] # other considered strategies
|
||||
reason: str
|
||||
|
||||
|
||||
class StrategySelector:
|
||||
"""
|
||||
Selects the best strategy for current market conditions.
|
||||
|
||||
Flow:
|
||||
1. Classify current market regime
|
||||
2. Look up performance matrix for regime
|
||||
3. Select strategy with highest regime-specific score
|
||||
4. Apply confidence threshold (fallback to default if low confidence)
|
||||
5. Return selection with alternatives
|
||||
|
||||
Adversarial testing:
|
||||
- Each strategy is tested across ALL regimes
|
||||
- Performance matrix tracks regime-specific scores
|
||||
- Selector floats best strategy to top per regime
|
||||
- Fallback to baseline if no strategy has enough data
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
classifier: Optional[RegimeClassifier] = None,
|
||||
matrix: Optional[PerformanceMatrix] = None,
|
||||
default_strategy_id: str = "baseline",
|
||||
min_episodes_for_selection: int = 3,
|
||||
confidence_threshold: float = 0.5,
|
||||
) -> None:
|
||||
self._classifier = classifier or RegimeClassifier()
|
||||
self._matrix = matrix or PerformanceMatrix()
|
||||
self._default_strategy_id = default_strategy_id
|
||||
self._min_episodes = min_episodes_for_selection
|
||||
self._confidence_threshold = confidence_threshold
|
||||
self._selection_history: list[SelectionResult] = []
|
||||
|
||||
def select(
|
||||
self,
|
||||
state: MarketWorldState,
|
||||
available_strategies: Dict[str, FulfilmentPolicyParams],
|
||||
) -> SelectionResult:
|
||||
"""
|
||||
Select the best strategy for current market conditions.
|
||||
|
||||
Args:
|
||||
state: current market state
|
||||
available_strategies: {strategy_id: params} mapping
|
||||
|
||||
Returns:
|
||||
SelectionResult with chosen strategy and reasoning
|
||||
"""
|
||||
# 1. Classify current regime
|
||||
regime = self._classifier.classify(state)
|
||||
|
||||
# 2. Look up performance matrix
|
||||
best_id = self._matrix.get_best(regime, exclude={self._default_strategy_id})
|
||||
|
||||
# 3. Check if best strategy has enough data
|
||||
if best_id and best_id in available_strategies:
|
||||
scores = self._matrix.get_scores_for_regime(regime)
|
||||
best_score = next((s for s in scores if s.strategy_id == best_id), None)
|
||||
|
||||
if best_score and best_score.episodes >= self._min_episodes:
|
||||
# 4. Check confidence
|
||||
all_scores = [s.score for s in scores]
|
||||
if len(all_scores) > 1:
|
||||
score_range = max(all_scores) - min(all_scores)
|
||||
confidence = min(1.0, score_range / max(max(all_scores), 1e-9))
|
||||
else:
|
||||
confidence = 0.5
|
||||
|
||||
if confidence >= self._confidence_threshold:
|
||||
result = SelectionResult(
|
||||
strategy_id=best_id,
|
||||
regime=regime,
|
||||
score=best_score.score,
|
||||
confidence=confidence,
|
||||
alternatives=tuple(s.strategy_id for s in scores[:3] if s.strategy_id != best_id),
|
||||
reason=f"best_for_{regime.value}",
|
||||
)
|
||||
self._selection_history.append(result)
|
||||
return result
|
||||
|
||||
# 5. Fallback to default
|
||||
result = SelectionResult(
|
||||
strategy_id=self._default_strategy_id,
|
||||
regime=regime,
|
||||
score=0.0,
|
||||
confidence=0.0,
|
||||
alternatives=(),
|
||||
reason="fallback_default",
|
||||
)
|
||||
self._selection_history.append(result)
|
||||
return result
|
||||
|
||||
def record_outcome(
|
||||
self,
|
||||
strategy_id: str,
|
||||
regime: MarketRegime,
|
||||
score: float,
|
||||
pnl_bps: float = 0.0,
|
||||
drawdown_bps: float = 0.0,
|
||||
adverse_fill_ratio: float = 0.0,
|
||||
) -> None:
|
||||
"""Record a strategy's outcome for learning."""
|
||||
self._matrix.record(
|
||||
strategy_id=strategy_id,
|
||||
regime=regime,
|
||||
score=score,
|
||||
pnl_bps=pnl_bps,
|
||||
drawdown_bps=drawdown_bps,
|
||||
adverse_fill_ratio=adverse_fill_ratio,
|
||||
)
|
||||
|
||||
@property
|
||||
def matrix(self) -> PerformanceMatrix:
|
||||
return self._matrix
|
||||
|
||||
@property
|
||||
def classifier(self) -> RegimeClassifier:
|
||||
return self._classifier
|
||||
|
||||
@property
|
||||
def selection_history(self) -> list[SelectionResult]:
|
||||
return list(self._selection_history)
|
||||
|
||||
@property
|
||||
def total_selections(self) -> int:
|
||||
return len(self._selection_history)
|
||||
Reference in New Issue
Block a user