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:
Codex
2026-07-11 10:33:56 +02:00
parent 6fb55dadcb
commit ef2f8e8827
4 changed files with 909 additions and 0 deletions

View File

@@ -0,0 +1 @@
from malkhut.training.cma_trainer import CMAESTrainer, CMAParameterCodec

View File

@@ -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

View File

@@ -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)),
)

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