""" Continuous Training Pipeline — runs the full training cycle indefinitely. Unlike the bounded pipeline, this one: - Runs forever until shutdown signal - Logs metrics every N generations - Checkpoints state periodically - Handles graceful shutdown via signal - Adapts strategy pool continuously """ from __future__ import annotations import json import os import signal import sys import time from dataclasses import dataclass, field from typing import Any, Dict, List, Optional from malkhut.state import FulfilmentPolicyParams from malkhut.training.pipeline import TrainingPipeline, PipelineConfig from malkhut.training.generator import StrategyGenerator, GeneratorConfig from malkhut.training.registry import PolicyRegistry from malkhut.training.cma_trainer import ScenarioFactory from malkhut.storage.ch_store import MalkhutCHStore @dataclass class ContinuousConfig: """Configuration for continuous training.""" checkpoint_interval_s: int = 300 # checkpoint every 5 minutes metrics_interval_s: int = 60 # log metrics every minute max_generations_per_cycle: int = 10 # generations per training cycle max_evals_per_generation: int = 10 strategy_pool_max: int = 50 log_path: str = "continuous_training.log" class ContinuousTrainingPipeline: """ Continuous training pipeline that runs indefinitely. Cycles through: 1. Training pipeline (CMA-ES) 2. Strategy generator (genetic programming) 3. Planner diversity testing 4. Metrics logging 5. Checkpointing """ def __init__(self, config: Optional[ContinuousConfig] = None) -> None: self.config = config or ContinuousConfig() self._running = False self._shutdown_event = threading.Event() if 'threading' in dir() else None # Components self._store = MalkhutCHStore() self._store.ensure_tables() self._registry = PolicyRegistry(store=self._store) self._scenario_factory = ScenarioFactory() # Metrics self._cycle_count = 0 self._total_evals = 0 self._total_strategies = 0 self._best_score = -float("inf") self._start_time = time.time() self._last_checkpoint = time.time() self._last_metrics = time.time() def run(self) -> None: """Run the continuous training loop.""" self._running = True print("=" * 70) print("MALKHUT CONTINUOUS TRAINING PIPELINE") print(f"Config: checkpoint={self.config.checkpoint_interval_s}s, " f"metrics={self.config.metrics_interval_s}s, " f"generations/cycle={self.config.max_generations_per_cycle}") print("=" * 70) # Register signal handler for graceful shutdown def signal_handler(sig, frame): print("\nShutdown signal received...") self._running = False signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGTERM, signal_handler) try: while self._running: self._run_cycle() self._cycle_count += 1 # Periodic metrics if time.time() - self._last_metrics >= self.config.metrics_interval_s: self._log_metrics() self._last_metrics = time.time() # Periodic checkpoint if time.time() - self._last_checkpoint >= self.config.checkpoint_interval_s: self._checkpoint() self._last_checkpoint = time.time() except KeyboardInterrupt: print("\nInterrupted by user") finally: self._running = False self._log_final_metrics() def _run_cycle(self) -> None: """Run one training cycle.""" # 1. Training pipeline pipeline_config = PipelineConfig( max_generations=self.config.max_generations_per_cycle, max_evals_per_generation=self.config.max_evals_per_generation, max_time_s=300, # 5 minutes per cycle auto_promote=True, ) pipeline = TrainingPipeline( config=pipeline_config, registry=self._registry, log_path=self.config.log_path, ) result = pipeline.run( incumbent=_baseline(), symbols=("BTCUSDT",), ) self._total_evals += result.total_evals # Track improvement if result.best_score > self._best_score: improvement = result.best_score - self._best_score self._best_score = result.best_score print(f" Cycle {self._cycle_count}: improvement +{improvement:.2f} (best={self._best_score:.2f})") # 2. Strategy generator gen_config = GeneratorConfig( population_size=10, generations=2, tournament_size=3, elitism_count=1, ) generator = StrategyGenerator(config=gen_config, registry=self._registry) scenarios = self._scenario_factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3) gen_population = generator.evolve(_baseline(), scenarios) gen_count = len([g for g in gen_population if g.generation > 0]) self._total_strategies += gen_count for genome in gen_population: if genome.generation > 0: generator.add_to_pool(genome) def _log_metrics(self) -> None: """Log current metrics.""" elapsed = time.time() - self._start_time print(f" [{elapsed:.0f}s] Cycle={self._cycle_count} " f"Evals={self._total_evals} Strategies={self._total_strategies} " f"Best={self._best_score:.2f}") def _checkpoint(self) -> None: """Checkpoint state to disk.""" checkpoint = { "timestamp": time.time(), "cycle_count": self._cycle_count, "total_evals": self._total_evals, "total_strategies": self._total_strategies, "best_score": self._best_score, "registry_records": self._registry.record_count, } with open("smoke_checkpoint.json", "w") as f: json.dump(checkpoint, f, indent=2) def _log_final_metrics(self) -> None: """Log final metrics.""" duration = time.time() - self._start_time print() print("=" * 70) print("CONTINUOUS TRAINING FINAL METRICS") print("=" * 70) print(f"Duration: {duration:.1f}s ({duration/60:.1f} min)") print(f"Cycles completed: {self._cycle_count}") print(f"Total evals: {self._total_evals}") print(f"Total strategies: {self._total_strategies}") print(f"Best score: {self._best_score:.2f}") print(f"Registry records: {self._registry.record_count}") print("=" * 70) def _baseline() -> FulfilmentPolicyParams: return FulfilmentPolicyParams( version="baseline", ucb_c=1.414, max_sims=256, max_depth=3, rollout_depth=3, root_temperature=0.5, min_root_entropy=0.25, quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.1, 0.25, 0.5), passive_ttl_ms=200, aggressive_ttl_ms=50, maker_edge_min_bps=0.5, cross_spread_edge_min_bps=5.0, adverse_toxicity_cancel_threshold=0.5, queue_churn_cancel_threshold=0.5, mae_tail_cut_bps=50.0, mfe_giveback_cut_fraction=0.5, max_time_in_loss_s=300.0, failed_recovery_cut_count=3, recovery_velocity_min_bps_per_s=0.0, max_symbol_notional_fraction=0.20, max_single_order_notional_fraction=0.05, reduce_when_global_up_fraction=0.30, session_profit_lock_fraction=0.02, w_expected_pnl=1.0, w_fill_probability=0.5, w_adverse_selection=2.0, w_queue_priority=0.5, w_inventory_risk=1.5, w_tail_loss=5.0, w_fee_quality=0.5, w_time_decay=0.3, w_policy_entropy=0.5, robust_tail_weight=2.0, toxic_counterparty_weight=3.0, low_liquidity_weight=2.0, latency_stress_weight=1.0, ) if __name__ == "__main__": import threading pipeline = ContinuousTrainingPipeline() pipeline.run()