212 lines
7.9 KiB
Python
212 lines
7.9 KiB
Python
|
|
"""
|
||
|
|
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()
|