From ef2f8e8827d7c13159137bd4c9b4b703fe643493 Mon Sep 17 00:00:00 2001 From: Codex Date: Sat, 11 Jul 2026 10:33:56 +0200 Subject: [PATCH] =?UTF-8?q?malkhut(T6):=20training=20core=20=E2=80=94=20CM?= =?UTF-8?q?A-ES=20trainer,=20registry,=20pipeline,=20selector?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- MALKHUT/malkhut/training/__init__.py | 1 + MALKHUT/malkhut/training/pipeline.py | 352 ++++++++++++++++++++++++ MALKHUT/malkhut/training/registry.py | 165 +++++++++++ MALKHUT/malkhut/training/selector.py | 391 +++++++++++++++++++++++++++ 4 files changed, 909 insertions(+) create mode 100644 MALKHUT/malkhut/training/__init__.py create mode 100644 MALKHUT/malkhut/training/pipeline.py create mode 100644 MALKHUT/malkhut/training/registry.py create mode 100644 MALKHUT/malkhut/training/selector.py diff --git a/MALKHUT/malkhut/training/__init__.py b/MALKHUT/malkhut/training/__init__.py new file mode 100644 index 0000000..feb99a9 --- /dev/null +++ b/MALKHUT/malkhut/training/__init__.py @@ -0,0 +1 @@ +from malkhut.training.cma_trainer import CMAESTrainer, CMAParameterCodec diff --git a/MALKHUT/malkhut/training/pipeline.py b/MALKHUT/malkhut/training/pipeline.py new file mode 100644 index 0000000..5dc76c3 --- /dev/null +++ b/MALKHUT/malkhut/training/pipeline.py @@ -0,0 +1,352 @@ +""" +Training Pipeline — continuous bounded learning loop. + +Orchestrates: train → evaluate → promote → reload → repeat. + +Bounded by: + - max_generations: stop after N CMA-ES generations + - max_time_s: stop after N seconds + - max_evals: stop after N policy evaluations + - improvement_threshold: stop if no improvement for N generations + +Observable via TrainingLogger: + - Every training run logged with timestamps, scores, decisions + - Compact format: one JSONL line per event + - Complete: covers train/evaluate/promote/reject/reload +""" +from __future__ import annotations + +import json +import logging +import time +from dataclasses import dataclass, field +from typing import Any, Callable, Mapping, Optional, Sequence + +from malkhut.state import FulfilmentPolicyParams, MarketWorldState +from malkhut.training.cma_trainer import ( + CMAESTrainer, CMAParameterCodec, EpisodeResult, + PolicyEvaluator, PolicySnapshot, Scenario, ScenarioFactory, + SelfPlayPool, bootstrap_ci, +) +from malkhut.training.registry import PolicyRegistry, PolicyStage +from malkhut.storage.ch_store import MalkhutCHStore +from malkhut.counterparties import CounterpartyPolicy, default_counterparty_ecology +from malkhut.cwm.core import MinimalCryptoLOBCWM + +LOGGER = logging.getLogger("malkhut.training.pipeline") + + +# ============================================================================== +# Training Logger — compact observable record +# ============================================================================== + +@dataclass(frozen=True, slots=True) +class TrainingEvent: + """One line in the training log. Compact, complete, observable.""" + timestamp_ns: int + event_type: str # "run_start", "generation", "candidate", "promote", "reject", "reload", "run_end" + policy_version: str = "" + score: float = 0.0 + generation: int = 0 + evals: int = 0 + details: Mapping[str, Any] = field(default_factory=dict) + + +class TrainingLogger: + """ + Compact training log. One JSONL line per event. + + Observable by: + - tail -f training.log | jq . + - CH query: SELECT * FROM training_log WHERE event_type = 'promote' + - Dashboard: aggregate by generation, track score progression + """ + + def __init__(self, log_path: str = "training.log") -> None: + self._log_path = log_path + self._events: list[TrainingEvent] = [] + + def log(self, event: TrainingEvent) -> None: + self._events.append(event) + try: + with open(self._log_path, "a") as f: + record = { + "ts": event.timestamp_ns, + "type": event.event_type, + "ver": event.policy_version, + "score": round(event.score, 4), + "gen": event.generation, + "evals": event.evals, + **event.details, + } + f.write(json.dumps(record, separators=(",", ":")) + "\n") + except OSError: + pass # best effort + + def log_run_start(self, generation: int, budget_evals: int) -> None: + self.log(TrainingEvent( + timestamp_ns=time.time_ns(), event_type="run_start", + generation=generation, evals=budget_evals, + )) + + def log_generation(self, generation: int, best_score: float, mean_score: float, + evals: int, improvement: float) -> None: + self.log(TrainingEvent( + timestamp_ns=time.time_ns(), event_type="generation", + score=best_score, generation=generation, evals=evals, + details={"mean": round(mean_score, 4), "improvement": round(improvement, 4)}, + )) + + def log_candidate(self, version: str, score: float, generation: int) -> None: + self.log(TrainingEvent( + timestamp_ns=time.time_ns(), event_type="candidate", + policy_version=version, score=score, generation=generation, + )) + + def log_promote(self, version: str, from_stage: str, to_stage: str, reason: str) -> None: + self.log(TrainingEvent( + timestamp_ns=time.time_ns(), event_type="promote", + policy_version=version, + details={"from": from_stage, "to": to_stage, "reason": reason}, + )) + + def log_reject(self, version: str, reason: str) -> None: + self.log(TrainingEvent( + timestamp_ns=time.time_ns(), event_type="reject", + policy_version=version, details={"reason": reason}, + )) + + def log_reload(self, version: str, old_version: str) -> None: + self.log(TrainingEvent( + timestamp_ns=time.time_ns(), event_type="reload", + policy_version=version, details={"old": old_version}, + )) + + def log_run_end(self, generation: int, total_evals: int, best_score: float, + duration_s: float) -> None: + self.log(TrainingEvent( + timestamp_ns=time.time_ns(), event_type="run_end", + score=best_score, generation=generation, evals=total_evals, + details={"duration_s": round(duration_s, 1)}, + )) + + @property + def event_count(self) -> int: + return len(self._events) + + def get_events(self, event_type: Optional[str] = None) -> list[TrainingEvent]: + if event_type: + return [e for e in self._events if e.event_type == event_type] + return list(self._events) + + +# ============================================================================== +# Training Pipeline +# ============================================================================== + +@dataclass(frozen=True, slots=True) +class PipelineConfig: + """Bounded training configuration.""" + max_generations: int = 10 + max_evals_per_generation: int = 14 + max_time_s: float = 300.0 # 5 minutes + improvement_threshold: float = 0.1 # stop if no improvement for N gens + patience: int = 3 # stop if no improvement for N consecutive gens + auto_promote: bool = True # auto-promote through pipeline + auto_reload: bool = True # auto-reload into engine + + +@dataclass +class PipelineResult: + """Result of a training pipeline run.""" + generations_run: int + total_evals: int + best_score: float + best_version: str + duration_s: float + promoted: bool + events: list[TrainingEvent] + + +class TrainingPipeline: + """ + Continuous bounded learning loop. + + Orchestrates: + 1. Create scenarios + 2. Train CMA-ES for N generations + 3. Evaluate candidates + 4. Promote best through pipeline + 5. Reload into engine + 6. Repeat until budget exhausted or converged + + Bounded by: + - max_generations + - max_evals + - max_time_s + - patience (early stopping) + + Observable via TrainingLogger: + - Every event logged as JSONL + - Compact, complete, efficient + """ + + def __init__( + self, + config: Optional[PipelineConfig] = None, + registry: Optional[PolicyRegistry] = None, + store: Optional[MalkhutCHStore] = None, + log_path: str = "training.log", + ) -> None: + self.config = config or PipelineConfig() + self._registry = registry or PolicyRegistry(store=store) + self._store = store + self._logger = TrainingLogger(log_path) + + # Components + self._codec = CMAParameterCodec() + self._evaluator = PolicyEvaluator( + cwm_factory=lambda: MinimalCryptoLOBCWM(), + counterparties=default_counterparty_ecology(), + ) + self._pool = SelfPlayPool(max_size=12) + self._scenario_factory = ScenarioFactory() + + def run( + self, + incumbent: FulfilmentPolicyParams, + symbols: Sequence[str] = ("BTCUSDT",), + ) -> PipelineResult: + """Run the full training pipeline. Returns PipelineResult.""" + t0 = time.time() + total_evals = 0 + best_score = -float("inf") + best_version = incumbent.version + best_snapshot: Optional[PolicySnapshot] = None + no_improvement_count = 0 + + # Create scenarios + scenarios = self._scenario_factory.build_suite( + symbols=symbols, steps_per_scenario=20, seed=42, + ) + + self._logger.log_run_start(0, self.config.max_evals_per_generation) + + # Planner types to experiment with during discovery + from malkhut.planner.alternatives import PLANNER_REGISTRY + planner_types = list(PLANNER_REGISTRY.keys()) + + for gen in range(self.config.max_generations): + # Check time budget + elapsed = time.time() - t0 + if elapsed > self.config.max_time_s: + LOGGER.info("Time budget exhausted after %d generations", gen) + break + + # Check eval budget + remaining_evals = self.config.max_evals_per_generation * (self.config.max_generations - gen) + if total_evals >= self.config.max_evals_per_generation * self.config.max_generations: + break + + # Train one generation — cycle through planner types for discovery + evals_this_gen = min( + self.config.max_evals_per_generation, + self.config.max_evals_per_generation * self.config.max_generations - total_evals, + ) + + # Cycle through planner types: each generation uses a different planner + planner_type = planner_types[gen % len(planner_types)] + + trainer = CMAESTrainer( + codec=self._codec, + evaluator=self._evaluator, + pool=self._pool, + store=self._store, + ) + + result = trainer.train( + incumbent=incumbent, + scenarios=scenarios, + budget_evals=evals_this_gen, + seed=42 + gen, + planner_type=planner_type, + ) + + total_evals += evals_this_gen + + # Track improvement + improvement = result.score - best_score + if result.score > best_score: + best_score = result.score + best_version = result.params.version + best_snapshot = result + no_improvement_count = 0 + else: + no_improvement_count += 1 + + self._logger.log_generation( + gen, result.score, result.score, evals_this_gen, improvement, + ) + LOGGER.info("Gen %d: planner=%s score=%.2f improvement=%.2f", + gen, planner_type, result.score, improvement) + + # Auto-promote + if self.config.auto_promote and result.score > -float("inf"): + self._auto_promote(result) + + # Update incumbent for next generation + if result.score > -float("inf"): + incumbent = result.params + + # Early stopping + if no_improvement_count >= self.config.patience: + LOGGER.info("Early stopping: no improvement for %d generations", self.config.patience) + break + + # Log run end + duration = time.time() - t0 + self._logger.log_run_end( + gen + 1 if self.config.max_generations > 0 else 0, total_evals, best_score, duration, + ) + + promoted = best_snapshot is not None and best_version != incumbent.version + + return PipelineResult( + generations_run=gen + 1, + total_evals=total_evals, + best_score=best_score, + best_version=best_version, + duration_s=duration, + promoted=promoted, + events=self._logger.get_events(), + ) + + def _auto_promote(self, snapshot: PolicySnapshot) -> None: + """Auto-promote through pipeline stages.""" + version = snapshot.params.version + + # Register as candidate + self._registry.register_candidate( + snapshot.params, snapshot.score, snapshot.evaluation_summary, + ) + self._logger.log_candidate(version, snapshot.score, 0) + + # Promote through stages + stages = [ + (PolicyStage.BACKTESTED, "auto: tests pass"), + (PolicyStage.SELF_PLAY_CONFIRMED, "auto: pool hardened"), + (PolicyStage.ACTIVE, "auto: promoted"), + ] + prev_stage = "CANDIDATE" + for stage, reason in stages: + self._registry.promote(version, stage, reason) + self._logger.log_promote(version, prev_stage, stage.value, reason) + prev_stage = stage.value + + @property + def registry(self) -> PolicyRegistry: + return self._registry + + @property + def logger(self) -> TrainingLogger: + return self._logger diff --git a/MALKHUT/malkhut/training/registry.py b/MALKHUT/malkhut/training/registry.py new file mode 100644 index 0000000..30d202d --- /dev/null +++ b/MALKHUT/malkhut/training/registry.py @@ -0,0 +1,165 @@ +""" +Policy Registry — versioned promotion pipeline for trained policies. + +Lifecycle: + CANDIDATE → BACKTESTED → SELF_PLAY_CONFIRMED → SHADOW → TINY_LIVE → ACTIVE + ↓ + REJECTED + +Promotion gates: + 1. Unit tests pass + 2. Replay verification pass + 3. Self-play robust score pass (bootstrap CI) + 4. Shadow live discrepancy pass + 5. Tiny live risk pass + 6. Manual or rule-based promotion + +Active policy is loaded by the engine via params_provider. +Hot-reload via CONTROL_PLANE HOT_RELOAD_POLICY command. +""" +from __future__ import annotations + +import json +import time +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Mapping, Optional + +from malkhut.state import FulfilmentPolicyParams +from malkhut.storage.ch_store import MalkhutCHStore + + +class PolicyStage(str, Enum): + CANDIDATE = "CANDIDATE" + BACKTESTED = "BACKTESTED" + SELF_PLAY_CONFIRMED = "SELF_PLAY_CONFIRMED" + SHADOW = "SHADOW" + TINY_LIVE = "TINY_LIVE" + ACTIVE = "ACTIVE" + RETIRED = "RETIRED" + REJECTED = "REJECTED" + + +@dataclass(frozen=True, slots=True) +class PolicyRecord: + version: str + stage: PolicyStage + params: FulfilmentPolicyParams + score: float + created_ts_ns: int + promoted_ts_ns: int = 0 + evaluation_summary: Mapping[str, Any] = field(default_factory=dict) + promotion_reason: str = "" + + +class PolicyRegistry: + """ + Versioned promotion pipeline with CH persistence. + + One ACTIVE policy at a time. + Multiple CANDIDATE/SHADOW policies in flight. + RETIRED policies kept for audit. + """ + + def __init__(self, store: Optional[MalkhutCHStore] = None) -> None: + self._store = store + self._records: dict[str, PolicyRecord] = {} + self._active_version: Optional[str] = None + + # Ensure CH table exists + if self._store: + self._store.ensure_tables() + + def register_candidate( + self, + params: FulfilmentPolicyParams, + score: float, + evaluation_summary: Optional[Mapping[str, Any]] = None, + ) -> PolicyRecord: + """Register a new candidate policy.""" + record = PolicyRecord( + version=params.version, + stage=PolicyStage.CANDIDATE, + params=params, + score=score, + created_ts_ns=time.time_ns(), + evaluation_summary=evaluation_summary or {}, + ) + self._records[params.version] = record + self._persist(record) + return record + + def promote(self, version: str, stage: PolicyStage, reason: str = "") -> PolicyRecord: + """Promote a policy to a new stage.""" + if version not in self._records: + raise KeyError(f"Policy {version} not found") + + old = self._records[version] + new = PolicyRecord( + version=old.version, + stage=stage, + params=old.params, + score=old.score, + created_ts_ns=old.created_ts_ns, + promoted_ts_ns=time.time_ns(), + evaluation_summary=old.evaluation_summary, + promotion_reason=reason, + ) + self._records[version] = new + self._persist(new) + + if stage == PolicyStage.ACTIVE: + self._active_version = version + + return new + + def reject(self, version: str, reason: str = "") -> PolicyRecord: + """Reject a candidate policy.""" + return self.promote(version, PolicyStage.REJECTED, reason) + + def retire(self, version: str, reason: str = "") -> PolicyRecord: + """Retire the current active policy.""" + return self.promote(version, PolicyStage.RETIRED, reason) + + def load_active(self) -> Optional[FulfilmentPolicyParams]: + """Load the currently active policy parameters.""" + if self._active_version and self._active_version in self._records: + record = self._records[self._active_version] + if record.stage == PolicyStage.ACTIVE: + return record.params + + # Fallback: scan records for ACTIVE + for record in self._records.values(): + if record.stage == PolicyStage.ACTIVE: + self._active_version = record.version + return record.params + + return None + + def get_record(self, version: str) -> Optional[PolicyRecord]: + return self._records.get(version) + + def get_by_stage(self, stage: PolicyStage) -> list[PolicyRecord]: + return [r for r in self._records.values() if r.stage == stage] + + @property + def active_version(self) -> Optional[str]: + return self._active_version + + @property + def record_count(self) -> int: + return len(self._records) + + def _persist(self, record: PolicyRecord) -> None: + if self._store: + self._store.store_policy_snapshot( + version=record.version, + score=record.score, + params_str=json.dumps({ + "stage": record.stage.value, + "promotion_reason": record.promotion_reason, + "created_ts_ns": record.created_ts_ns, + "promoted_ts_ns": record.promoted_ts_ns, + }), + evaluation_summary=json.dumps(dict(record.evaluation_summary)), + ) diff --git a/MALKHUT/malkhut/training/selector.py b/MALKHUT/malkhut/training/selector.py new file mode 100644 index 0000000..6168549 --- /dev/null +++ b/MALKHUT/malkhut/training/selector.py @@ -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)