256 lines
12 KiB
Python
256 lines
12 KiB
Python
|
|
"""
|
||
|
|
Tests for training pipeline and logging.
|
||
|
|
|
||
|
|
Verifies:
|
||
|
|
- Pipeline runs bounded iterations
|
||
|
|
- Early stopping on convergence
|
||
|
|
- Time budget respected
|
||
|
|
- Eval budget respected
|
||
|
|
- Logger produces compact JSONL
|
||
|
|
- Events are complete and observable
|
||
|
|
- Full pipeline flow: train → promote → reload
|
||
|
|
"""
|
||
|
|
import json
|
||
|
|
import os
|
||
|
|
import tempfile
|
||
|
|
import time
|
||
|
|
import pytest
|
||
|
|
from malkhut.state import FulfilmentPolicyParams
|
||
|
|
from malkhut.training.pipeline import (
|
||
|
|
TrainingPipeline, PipelineConfig, PipelineResult,
|
||
|
|
TrainingLogger, TrainingEvent,
|
||
|
|
)
|
||
|
|
from malkhut.training.registry import PolicyRegistry, PolicyStage
|
||
|
|
|
||
|
|
|
||
|
|
def _baseline(**kw):
|
||
|
|
d = dict(
|
||
|
|
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,
|
||
|
|
)
|
||
|
|
d.update(kw)
|
||
|
|
return FulfilmentPolicyParams(**d)
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# 1. TRAINING LOGGER
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestTrainingLogger:
|
||
|
|
def test_log_creates_file(self):
|
||
|
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||
|
|
path = f.name
|
||
|
|
try:
|
||
|
|
logger = TrainingLogger(log_path=path)
|
||
|
|
logger.log(TrainingEvent(
|
||
|
|
timestamp_ns=time.time_ns(), event_type="test",
|
||
|
|
policy_version="v1", score=10.0,
|
||
|
|
))
|
||
|
|
assert os.path.exists(path)
|
||
|
|
with open(path) as f:
|
||
|
|
lines = f.readlines()
|
||
|
|
assert len(lines) == 1
|
||
|
|
record = json.loads(lines[0])
|
||
|
|
assert record["type"] == "test"
|
||
|
|
assert record["ver"] == "v1"
|
||
|
|
assert record["score"] == 10.0
|
||
|
|
finally:
|
||
|
|
os.unlink(path)
|
||
|
|
|
||
|
|
def test_log_run_start(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log_run_start(generation=0, budget_evals=14)
|
||
|
|
assert logger.event_count == 1
|
||
|
|
assert logger.get_events()[0].event_type == "run_start"
|
||
|
|
|
||
|
|
def test_log_generation(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log_generation(generation=1, best_score=5.0, mean_score=3.0, evals=14, improvement=2.0)
|
||
|
|
assert logger.event_count == 1
|
||
|
|
e = logger.get_events()[0]
|
||
|
|
assert e.score == 5.0
|
||
|
|
assert e.details["mean"] == 3.0
|
||
|
|
|
||
|
|
def test_log_candidate(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log_candidate("v1", 10.0, 1)
|
||
|
|
assert logger.get_events()[0].policy_version == "v1"
|
||
|
|
|
||
|
|
def test_log_promote(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log_promote("v1", "CANDIDATE", "ACTIVE", "promoted")
|
||
|
|
e = logger.get_events()[0]
|
||
|
|
assert e.event_type == "promote"
|
||
|
|
assert e.details["from"] == "CANDIDATE"
|
||
|
|
assert e.details["to"] == "ACTIVE"
|
||
|
|
|
||
|
|
def test_log_reject(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log_reject("v1", "tail risk")
|
||
|
|
assert logger.get_events()[0].event_type == "reject"
|
||
|
|
|
||
|
|
def test_log_reload(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log_reload("v2", "v1")
|
||
|
|
e = logger.get_events()[0]
|
||
|
|
assert e.details["old"] == "v1"
|
||
|
|
|
||
|
|
def test_log_run_end(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log_run_end(generation=5, total_evals=70, best_score=8.0, duration_s=120.0)
|
||
|
|
e = logger.get_events()[0]
|
||
|
|
assert e.details["duration_s"] == 120.0
|
||
|
|
|
||
|
|
def test_filter_by_event_type(self):
|
||
|
|
logger = TrainingLogger(log_path="/dev/null")
|
||
|
|
logger.log(TrainingEvent(timestamp_ns=1, event_type="start"))
|
||
|
|
logger.log(TrainingEvent(timestamp_ns=2, event_type="generation"))
|
||
|
|
logger.log(TrainingEvent(timestamp_ns=3, event_type="start"))
|
||
|
|
starts = logger.get_events(event_type="start")
|
||
|
|
assert len(starts) == 2
|
||
|
|
|
||
|
|
def test_jsonl_compact(self):
|
||
|
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||
|
|
path = f.name
|
||
|
|
try:
|
||
|
|
logger = TrainingLogger(log_path=path)
|
||
|
|
logger.log(TrainingEvent(
|
||
|
|
timestamp_ns=1234567890, event_type="test",
|
||
|
|
policy_version="v1", score=10.0, generation=1, evals=14,
|
||
|
|
details={"mean": 5.0},
|
||
|
|
))
|
||
|
|
with open(path) as f:
|
||
|
|
line = f.readline()
|
||
|
|
# Compact: no spaces after separators
|
||
|
|
assert ",\"score\":10.0" in line
|
||
|
|
assert "\"mean\":5.0" in line
|
||
|
|
finally:
|
||
|
|
os.unlink(path)
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# 2. PIPELINE CONFIG
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestPipelineConfig:
|
||
|
|
def test_default_config(self):
|
||
|
|
cfg = PipelineConfig()
|
||
|
|
assert cfg.max_generations == 10
|
||
|
|
assert cfg.max_evals_per_generation == 14
|
||
|
|
assert cfg.max_time_s == 300.0
|
||
|
|
assert cfg.patience == 3
|
||
|
|
|
||
|
|
def test_custom_config(self):
|
||
|
|
cfg = PipelineConfig(max_generations=5, patience=2)
|
||
|
|
assert cfg.max_generations == 5
|
||
|
|
assert cfg.patience == 2
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# 3. PIPELINE RUN
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestTrainingPipeline:
|
||
|
|
def test_pipeline_returns_result(self):
|
||
|
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
assert isinstance(result, PipelineResult)
|
||
|
|
assert result.generations_run >= 1
|
||
|
|
assert result.total_evals > 0
|
||
|
|
assert result.duration_s > 0
|
||
|
|
|
||
|
|
def test_pipeline_respects_max_generations(self):
|
||
|
|
config = PipelineConfig(max_generations=2, max_evals_per_generation=3, max_time_s=60)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
assert result.generations_run <= 2
|
||
|
|
|
||
|
|
def test_pipeline_respects_time_budget(self):
|
||
|
|
config = PipelineConfig(max_generations=100, max_evals_per_generation=3, max_time_s=2.0)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
# Should stop within reasonable time
|
||
|
|
assert result.generations_run <= 10 # not all 100
|
||
|
|
assert result.duration_s < 60.0
|
||
|
|
|
||
|
|
def test_pipeline_early_stopping(self):
|
||
|
|
config = PipelineConfig(max_generations=100, max_evals_per_generation=3, patience=1, max_time_s=60)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
assert result.generations_run <= 100
|
||
|
|
|
||
|
|
def test_pipeline_logs_events(self):
|
||
|
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
events = pipeline.logger.get_events()
|
||
|
|
assert len(events) > 0
|
||
|
|
event_types = [e.event_type for e in events]
|
||
|
|
assert "run_start" in event_types
|
||
|
|
assert "run_end" in event_types
|
||
|
|
assert "generation" in event_types
|
||
|
|
|
||
|
|
def test_pipeline_has_best_score(self):
|
||
|
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
assert isinstance(result.best_score, float)
|
||
|
|
|
||
|
|
def test_pipeline_registry_has_records(self):
|
||
|
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30, auto_promote=True)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
assert pipeline.registry.record_count > 0
|
||
|
|
|
||
|
|
def test_pipeline_result_events_match_logger(self):
|
||
|
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
assert len(result.events) == pipeline.logger.event_count
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# 4. FULL OBSERVABLE FLOW
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestFullObservableFlow:
|
||
|
|
def test_train_log_promote_activate(self):
|
||
|
|
"""Full observable flow: train → log → promote → activate → engine loads."""
|
||
|
|
from malkhut.engine import FulfilmentEngine
|
||
|
|
|
||
|
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||
|
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||
|
|
|
||
|
|
# Engine starts with default
|
||
|
|
engine = FulfilmentEngine(registry=pipeline.registry)
|
||
|
|
|
||
|
|
# Run pipeline
|
||
|
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||
|
|
|
||
|
|
# Verify logging happened
|
||
|
|
assert len(result.events) > 0
|
||
|
|
|
||
|
|
# Verify registry has promoted policies
|
||
|
|
active = pipeline.registry.load_active()
|
||
|
|
if active:
|
||
|
|
# Hot reload engine
|
||
|
|
engine.hot_reload_policy(active)
|
||
|
|
assert engine.params_provider().version == active.version
|
||
|
|
|
||
|
|
engine.close()
|