353 lines
12 KiB
Python
353 lines
12 KiB
Python
|
|
"""
|
||
|
|
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
|