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

123 lines
4.6 KiB
Python
Raw Normal View History

"""
CMA parameter codec — encode/decode roundtrip, bounds, type preservation.
"""
import pytest
from malkhut.training.cma_trainer import CMAParameterCodec
from malkhut.state import FulfilmentPolicyParams
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)
class TestCodecBounds:
def test_bounds_length_matches_specs(self):
codec = CMAParameterCodec()
lows, highs = codec.bounds()
assert len(lows) == len(codec.SPECS)
assert len(highs) == len(codec.SPECS)
def test_lows_less_than_highs(self):
codec = CMAParameterCodec()
lows, highs = codec.bounds()
for lo, hi in zip(lows, highs):
assert lo < hi
def test_initial_vector_midpoint(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
lows, highs = codec.bounds()
for i, (v, lo, hi) in enumerate(zip(x0, lows, highs)):
assert lo <= v <= hi
def test_initial_vector_length(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
assert len(x0) == len(codec.SPECS)
class TestCodecDecode:
def test_decode_returns_fulfilment_policy_params(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
p = codec.decode(x0, "v_test")
assert isinstance(p, FulfilmentPolicyParams)
def test_version_preserved(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
p = codec.decode(x0, "my_version")
assert p.version == "my_version"
def test_int_fields_are_integers(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
p = codec.decode(x0, "int_test")
assert isinstance(p.max_depth, int)
assert isinstance(p.passive_ttl_ms, int)
assert isinstance(p.failed_recovery_cut_count, int)
def test_float_fields_are_floats(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
p = codec.decode(x0, "float_test")
assert isinstance(p.ucb_c, float)
assert isinstance(p.mae_tail_cut_bps, float)
def test_bounds_clipping_above(self):
codec = CMAParameterCodec()
highs = [s.high for s in codec.SPECS]
x_over = [h + 10.0 for h in highs]
p = codec.decode(x_over, "over")
lows, highs_b = codec.bounds()
for i, spec in enumerate(codec.SPECS):
val = getattr(p, spec.name)
if spec.kind == "float":
assert val <= spec.high + 1e-9
def test_bounds_clipping_below(self):
codec = CMAParameterCodec()
lows = [s.low for s in codec.SPECS]
x_under = [l - 10.0 for l in lows]
p = codec.decode(x_under, "under")
for i, spec in enumerate(codec.SPECS):
val = getattr(p, spec.name)
if spec.kind == "float":
assert val >= spec.low - 1e-9
def test_decode_idempotent_at_midpoint(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
p1 = codec.decode(x0, "v1")
p2 = codec.decode(x0, "v2")
assert p1.ucb_c == p2.ucb_c
assert p1.max_depth == p2.max_depth
def test_different_vectors_different_params(self):
codec = CMAParameterCodec()
x0 = codec.initial_vector(_baseline())
x1 = list(x0)
x1[0] = x0[0] + 0.5 # ucb_c
p0 = codec.decode(x0, "a")
p1 = codec.decode(x1, "b")
assert p0.ucb_c != p1.ucb_c