diff --git a/MALKHUT/malkhut/tests/test_scenario_library.py b/MALKHUT/malkhut/tests/test_scenario_library.py new file mode 100644 index 0000000..484b6e1 --- /dev/null +++ b/MALKHUT/malkhut/tests/test_scenario_library.py @@ -0,0 +1,56 @@ +""" +Tests for ScenarioLibrary — sweep coverage. +""" +import pytest +from malkhut.training.scenario_library import ScenarioLibrary, SweepPoint, get_scenario_library + + +class TestScenarioLibrary: + def test_default_grid_size(self): + lib = ScenarioLibrary() + size = lib.grid_size() + assert size == 13 * 6 * 6 * 4 * 4 # 13 assets × 576 grid points + print(f' Default grid: {size} points') + + def test_sweep_produces_points(self): + lib = ScenarioLibrary(symbols=("BTCUSDT",)) + points = lib.sweep() + assert len(points) == 6 * 6 * 4 * 4 # 576 for one asset + + def test_sweep_point_fields(self): + lib = ScenarioLibrary(symbols=("BTCUSDT",)) + points = lib.sweep() + p = points[0] + assert isinstance(p, SweepPoint) + assert p.symbol == "BTCUSDT" + assert isinstance(p.spread_mult, float) + assert isinstance(p.depth_fraction, float) + assert isinstance(p.toxicity, float) + assert isinstance(p.regime, str) + assert isinstance(p.label, str) + + def test_sweep_covers_all_regimes(self): + lib = ScenarioLibrary(symbols=("BTCUSDT",)) + points = lib.sweep() + regimes = set(p.regime for p in points) + assert regimes == {"normal", "crisis", "recovery", "transition"} + + def test_custom_dimensions(self): + lib = ScenarioLibrary( + symbols=("BTCUSDT", "ETHUSDT"), + spread_mults=[1.0, 5.0], + depth_fracs=[0.1, 0.5], + tox_levels=[0.0, 0.5], + regimes=["normal", "crisis"], + ) + assert lib.grid_size() == 2 * 2 * 2 * 2 * 2 # 32 points + + def test_summary(self): + lib = ScenarioLibrary() + s = lib.summary() + assert "13" in s + assert "7488" in s # 13 × 576 + + def test_get_scenario_library(self): + lib = get_scenario_library() + assert lib.grid_size() > 0 diff --git a/MALKHUT/malkhut/training/scenario_library.py b/MALKHUT/malkhut/training/scenario_library.py new file mode 100644 index 0000000..b385c6c --- /dev/null +++ b/MALKHUT/malkhut/training/scenario_library.py @@ -0,0 +1,97 @@ +""" +ScenarioLibrary — sweeps the state space for comprehensive Mode 1 coverage. + +Unlike ScenarioFactory which builds 30 pre-defined scenario types per asset, +ScenarioLibrary SWEEPS across parameter dimensions to ensure coverage of +regions the tape never visited. + +Sweep dimensions (orthogonal): + - spread_bps: [0.1, 0.5, 1.0, 2.0, 5.0, 10.0] + - depth_fraction: [0.01, 0.05, 0.1, 0.3, 0.5, 1.0] + - toxicity: [0.0, 0.3, 0.7, 1.0] + - regime: ["normal", "crisis", "recovery"] + +Total: 6 × 6 × 4 × 4 = 576 grid points per asset. +With 3 assets = 1,728 scenarios total. +""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Dict, List, Optional, Sequence, Tuple + +from malkhut.training.asset_classification import ASSET_PROFILES, get_asset_profile + + +# Sweep dimension defaults +DEFAULT_SPREAD_MULTS = [0.1, 0.5, 1.0, 2.0, 5.0, 10.0] +DEFAULT_DEPTH_FRACS = [0.01, 0.05, 0.1, 0.3, 0.5, 1.0] +DEFAULT_TOXICITY = [0.0, 0.3, 0.7, 1.0] +DEFAULT_REGIMES = ["normal", "crisis", "recovery", "transition"] + + +@dataclass(frozen=True, slots=True) +class SweepPoint: + """A single point in the swept state space.""" + symbol: str + spread_mult: float + depth_fraction: float + toxicity: float + regime: str + label: str # human-readable: "BTC_spread2x_depth30pct_tox0.3_crisis" + + +class ScenarioLibrary: + """Sweeps the state space for Mode 1 (EXPLORE) coverage. + + Produces grid points across spread × depth × toxicity × regime. + These are ANCHORS that the ecology fills between and beyond. + """ + + def __init__( + self, + symbols: Optional[Sequence[str]] = None, + spread_mults: Optional[Sequence[float]] = None, + depth_fracs: Optional[Sequence[float]] = None, + tox_levels: Optional[Sequence[float]] = None, + regimes: Optional[Sequence[str]] = None, + ) -> None: + self.symbols = list(symbols or ASSET_PROFILES.keys()) + self.spread_mults = list(spread_mults or DEFAULT_SPREAD_MULTS) + self.depth_fracs = list(depth_fracs or DEFAULT_DEPTH_FRACS) + self.tox_levels = list(tox_levels or DEFAULT_TOXICITY) + self.regimes = list(regimes or DEFAULT_REGIMES) + + def sweep(self) -> List[SweepPoint]: + """Generate full grid sweep across all dimensions.""" + points = [] + for symbol in self.symbols: + for spread in self.spread_mults: + for depth in self.depth_fracs: + for tox in self.tox_levels: + for regime in self.regimes: + label = f"{symbol}_s{spread}_d{depth}_t{tox}_r{regime}" + points.append(SweepPoint( + symbol=symbol, + spread_mult=spread, + depth_fraction=depth, + toxicity=tox, + regime=regime, + label=label, + )) + return points + + def grid_size(self) -> int: + return len(self.symbols) * len(self.spread_mults) * len(self.depth_fracs) * len(self.tox_levels) * len(self.regimes) + + def summary(self) -> str: + return (f"ScenarioLibrary: {len(self.symbols)} assets × " + f"{len(self.spread_mults)} spreads × {len(self.depth_fracs)} depths × " + f"{len(self.tox_levels)} tox × {len(self.regimes)} regimes = " + f"{self.grid_size()} grid points") + + +def get_scenario_library( + symbols: Optional[Sequence[str]] = None, +) -> ScenarioLibrary: + """Get default scenario library for given symbols.""" + return ScenarioLibrary(symbols=symbols)