CWM (103): core mechanics, exhaustive edge cases, numba, exchange mechanics Replay (118): exhaustive verification, microstructure, trajectory Training (190): asset classification, phase0 extensive, pipeline, exhaustive DSL (102): v2 syntax, expanded, new features ASEx (33): validate-before-mutate, single-writer Planner (48): MCTS, alternatives, hooks Counterparties (19): 9 adversarial agent policies Clock (30): event-driven reactor BingX (28): venue adapter IPC (8): Zinc SHM Storage (9): ClickHouse Risk (4): hard invariants State (17): frozen dataclass invariants Integration: E2E, concurrency, sync/async seams, hypothesis, fuzz, adversarial
67 lines
2.8 KiB
Python
67 lines
2.8 KiB
Python
"""
|
|
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)
|