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:
1
MALKHUT/malkhut/training/__init__.py
Normal file
1
MALKHUT/malkhut/training/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
from malkhut.training.cma_trainer import CMAESTrainer, CMAParameterCodec
|
||||||
352
MALKHUT/malkhut/training/pipeline.py
Normal file
352
MALKHUT/malkhut/training/pipeline.py
Normal 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
|
||||||
165
MALKHUT/malkhut/training/registry.py
Normal file
165
MALKHUT/malkhut/training/registry.py
Normal 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)),
|
||||||
|
)
|
||||||
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