Files
sentiment-engine/MALKHUT/malkhut/tests/test_training.py

67 lines
2.8 KiB
Python
Raw Normal View History

"""
Unit tests: CMA parameter codec roundtrip.
"""
import pytest
from malkhut.training.cma_trainer import CMAParameterCodec
from malkhut.state import FulfilmentPolicyParams
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.10, 0.25, 0.50),
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,
)
class TestCMAParameterCodec:
def test_initial_vector_length(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
assert len(x0) == len(codec.SPECS)
def test_bounds_length(self):
codec = CMAParameterCodec()
lows, highs = codec.bounds()
assert len(lows) == len(codec.SPECS)
assert len(highs) == len(codec.SPECS)
assert all(l <= h for l, h in zip(lows, highs))
def test_decode_midpoint(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
params = codec.decode(x0, version="mid_test")
assert isinstance(params, FulfilmentPolicyParams)
assert params.version == "mid_test"
assert 0.2 <= params.ucb_c <= 3.0
def test_decode_clips_to_bounds(self):
codec = CMAParameterCodec()
lows, highs = codec.bounds()
x_over = [h + 1.0 for h in highs]
params = codec.decode(x_over, version="over")
for i, spec in enumerate(codec.SPECS):
val = getattr(params, spec.name)
assert spec.low <= val <= spec.high or isinstance(val, int)
def test_decode_int_fields_are_integers(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
params = codec.decode(x0, version="int_test")
assert isinstance(params.max_depth, int)
assert isinstance(params.passive_ttl_ms, int)
assert isinstance(params.failed_recovery_cut_count, int)