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:
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
|
||||
Reference in New Issue
Block a user