feat: add orderbook shadow training pipeline
This commit is contained in:
@@ -0,0 +1,175 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from crypto_spot_bot.models import Candle
|
||||
from crypto_spot_bot.orderbook_features import aggregate_orderbook_observations
|
||||
from crypto_spot_bot.shadow import shadow_gate_snapshot
|
||||
from crypto_spot_bot.storage import Storage
|
||||
from crypto_spot_bot.time_series import _feature_matrix
|
||||
from crypto_spot_bot.training_coordination import TrainingCoordinator
|
||||
|
||||
|
||||
def test_orderbook_aggregation_is_bucketed_and_rejects_sparse_hours() -> None:
|
||||
rows = [
|
||||
_observation(1_700_000_000_000, imbalance=0.6, spread=2.0, mid=100.0, micro=100.01),
|
||||
_observation(1_700_000_030_000, imbalance=0.2, spread=4.0, mid=100.0, micro=99.99),
|
||||
_observation(1_700_003_600_000, imbalance=-0.9, spread=8.0, mid=100.0, micro=100.02),
|
||||
]
|
||||
|
||||
features, manifest = aggregate_orderbook_observations(
|
||||
rows,
|
||||
interval="60",
|
||||
min_samples_per_bucket=2,
|
||||
)
|
||||
|
||||
assert manifest["BTCUSDT"]["covered_buckets"] == 1
|
||||
assert manifest["BTCUSDT"]["rejected_buckets"] == 1
|
||||
values = next(iter(features["BTCUSDT"].values()))
|
||||
assert values["l1_imbalance_mean"] == pytest.approx(0.4)
|
||||
assert values["l1_imbalance_std"] == pytest.approx(0.2)
|
||||
assert values["l1_spread_bps_mean"] == pytest.approx(3.0)
|
||||
assert values["l1_microprice_deviation_bps_mean"] == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_feature_matrix_uses_only_the_matching_closed_candle_bucket() -> None:
|
||||
candles = [
|
||||
Candle(timestamp=0, open=100, high=101, low=99, close=100, volume=1, turnover=100),
|
||||
Candle(timestamp=3_600_000, open=100, high=101, low=99, close=100, volume=1, turnover=100),
|
||||
]
|
||||
features = {
|
||||
"BTCUSDT": {
|
||||
0: {"l1_imbalance_mean": 0.25},
|
||||
3_600_000: {"l1_imbalance_mean": -0.75},
|
||||
7_200_000: {"l1_imbalance_mean": 0.99},
|
||||
}
|
||||
}
|
||||
|
||||
matrix = _feature_matrix(
|
||||
candles,
|
||||
["l1_imbalance_mean"],
|
||||
symbol="BTCUSDT",
|
||||
orderbook_features=features,
|
||||
)
|
||||
|
||||
assert matrix == [[0.25], [-0.75]]
|
||||
|
||||
|
||||
def test_shadow_gate_uses_only_settled_forward_predictions(tmp_path, monkeypatch) -> None:
|
||||
monkeypatch.setenv("SHADOW_GATE_MIN_SETTLED", "2")
|
||||
monkeypatch.setenv("SHADOW_GATE_MIN_ELIGIBLE", "2")
|
||||
monkeypatch.setenv("SHADOW_GATE_MIN_SYMBOLS", "1")
|
||||
monkeypatch.setenv("SHADOW_GATE_MIN_DIRECTION_ACCURACY", "0.5")
|
||||
monkeypatch.setenv("SHADOW_GATE_MAX_BRIER", "0.25")
|
||||
storage = Storage(tmp_path / "bot.sqlite3")
|
||||
model_sha = "a" * 64
|
||||
for timestamp, actual in ((1, 1.0), (2, 2.0)):
|
||||
storage.insert_shadow_prediction(
|
||||
model_sha256=model_sha,
|
||||
symbol="BTCUSDT",
|
||||
forecast_timestamp_ms=timestamp,
|
||||
horizon=1,
|
||||
reference_price=100.0,
|
||||
expected_return_percent=1.0,
|
||||
probability_up=0.8,
|
||||
eligible_signal=True,
|
||||
)
|
||||
row = storage.pending_shadow_predictions(
|
||||
model_sha256=model_sha,
|
||||
symbol="BTCUSDT",
|
||||
)[0]
|
||||
storage.settle_shadow_prediction(
|
||||
row["id"],
|
||||
actual_return_percent=actual,
|
||||
take_profit_first=True,
|
||||
)
|
||||
|
||||
gate = shadow_gate_snapshot(storage, model_sha)
|
||||
|
||||
assert gate["state"] == "passed"
|
||||
assert gate["settled_predictions"] == 2
|
||||
assert gate["eligible_predictions"] == 2
|
||||
assert gate["total_net_percent"] == pytest.approx(3.0)
|
||||
|
||||
|
||||
def test_shadow_bundle_does_not_replace_active_model(tmp_path) -> None:
|
||||
active = {"type": "active-model"}
|
||||
(tmp_path / "lstm_forecaster.json").write_text(json.dumps(active), encoding="utf-8")
|
||||
coordinator = TrainingCoordinator(tmp_path)
|
||||
job = coordinator.request_retrain({"source": "test"})["job"]
|
||||
lease_token = coordinator.claim({"worker_id": "worker-1"})["lease_token"]
|
||||
model = {
|
||||
"type": "pytorch_recurrent_forecaster",
|
||||
"symbols": {
|
||||
"BTCUSDT": {
|
||||
"model": "torch_gru",
|
||||
"lookback": 4,
|
||||
"input_size": 1,
|
||||
"hidden_size": 1,
|
||||
"state_dict": {"weight_ih_l0": [[0.0]]},
|
||||
"head_weight": [[0.0]],
|
||||
"head_bias": [0.0],
|
||||
}
|
||||
},
|
||||
}
|
||||
model_payload = (json.dumps(model) + "\n").encode()
|
||||
model_sha = hashlib.sha256(model_payload).hexdigest()
|
||||
artifacts = {
|
||||
"lstm_forecaster.shadow.json": model_payload,
|
||||
"torch_shadow_guard.json": (
|
||||
json.dumps({"accepted": True, "candidate_artifact_sha256": model_sha}) + "\n"
|
||||
).encode(),
|
||||
"torch_shadow_calibration.json": (
|
||||
json.dumps(
|
||||
{
|
||||
"artifact_sha256": model_sha,
|
||||
"validation": {
|
||||
"passed": True,
|
||||
"protocol": "untouched_model_holdout_with_threshold_walk_forward",
|
||||
},
|
||||
}
|
||||
)
|
||||
+ "\n"
|
||||
).encode(),
|
||||
}
|
||||
for name, payload in artifacts.items():
|
||||
coordinator.save_artifact_chunk(
|
||||
job["id"],
|
||||
{
|
||||
"name": name,
|
||||
"index": 0,
|
||||
"total": 1,
|
||||
"sha256": hashlib.sha256(payload).hexdigest(),
|
||||
"data_base64": base64.b64encode(payload).decode("ascii"),
|
||||
"lease_token": lease_token,
|
||||
},
|
||||
)
|
||||
|
||||
completed = coordinator.complete(
|
||||
job["id"],
|
||||
{"success": True, "summary": {"accepted": True, "deployment": "shadow"}, "lease_token": lease_token},
|
||||
)
|
||||
|
||||
assert completed["job"]["status"] == "completed"
|
||||
assert json.loads((tmp_path / "lstm_forecaster.json").read_text(encoding="utf-8")) == active
|
||||
assert json.loads((tmp_path / "lstm_forecaster.shadow.json").read_text(encoding="utf-8"))["type"] == "pytorch_recurrent_forecaster"
|
||||
|
||||
|
||||
def _observation(timestamp_ms: int, *, imbalance: float, spread: float, mid: float, micro: float) -> dict:
|
||||
return {
|
||||
"symbol": "BTCUSDT",
|
||||
"bid_price": mid - 0.01,
|
||||
"bid_size": 2.0,
|
||||
"ask_price": mid + 0.01,
|
||||
"ask_size": 1.0,
|
||||
"mid_price": mid,
|
||||
"microprice": micro,
|
||||
"spread_bps": spread,
|
||||
"imbalance": imbalance,
|
||||
"source_timestamp_ms": timestamp_ms,
|
||||
"created_at": "",
|
||||
}
|
||||
Reference in New Issue
Block a user