From 5d8ad1437ef28c5ef87d314e5048040939665e10 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D1=83=D1=80=D0=BD=D0=B0=D1=82=20=D0=90=D0=BD=D0=B4?= =?UTF-8?q?=D1=80=D0=B5=D0=B9?= Date: Wed, 15 Jul 2026 09:44:29 +0300 Subject: [PATCH] feat: add orderbook shadow training pipeline --- .env.example | 13 ++ README.md | 3 + android/TradeBotMonitor/app/build.gradle.kts | 4 +- .../kusoft/tradebotmonitor/MainActivity.kt | 32 +++ crypto_spot_bot/__init__.py | 2 +- crypto_spot_bot/bot.py | 116 +++++++++- crypto_spot_bot/dashboard.py | 52 ++++- crypto_spot_bot/market_data.py | 2 + crypto_spot_bot/orderbook_features.py | 180 +++++++++++++++ crypto_spot_bot/shadow.py | 94 ++++++++ crypto_spot_bot/storage.py | 184 ++++++++++++++- crypto_spot_bot/time_series.py | 39 +++- crypto_spot_bot/training_coordination.py | 146 +++++++++++- tests/test_orderbook_shadow.py | 175 +++++++++++++++ tools/calibrate_torch_thresholds.py | 34 ++- tools/run_torch_retrain.ps1 | 47 +++- tools/sync_orderbook_observations.py | 210 ++++++++++++++++++ tools/train_torch_recurrent_forecaster.py | 52 +++++ tools/windows_training_agent.py | 124 ++++++++++- 19 files changed, 1486 insertions(+), 23 deletions(-) create mode 100644 crypto_spot_bot/orderbook_features.py create mode 100644 crypto_spot_bot/shadow.py create mode 100644 tests/test_orderbook_shadow.py create mode 100644 tools/sync_orderbook_observations.py diff --git a/.env.example b/.env.example index b8b4813..41cc4ff 100644 --- a/.env.example +++ b/.env.example @@ -122,6 +122,19 @@ STORAGE_PRUNE_INTERVAL_SECONDS=3600 # Windows trainer keeps this final tail untouched by training and early stopping. TORCH_RETRAIN_HOLDOUT_WINDOW=1000 +TORCH_ORDERBOOK_DB=runtime/orderbook_observations.sqlite3 +TORCH_ORDERBOOK_MIN_SAMPLES_PER_BUCKET=20 +TORCH_ORDERBOOK_MIN_COVERED_BUCKETS=240 +TORCH_ORDERBOOK_MIN_SYMBOLS=2 + +# Forward-only gate for an offline-approved shadow model. Promotion remains an +# explicit authenticated API action after every check has passed. +SHADOW_GATE_MIN_SETTLED=300 +SHADOW_GATE_MIN_ELIGIBLE=30 +SHADOW_GATE_MIN_SYMBOLS=2 +SHADOW_GATE_MIN_PROFIT_FACTOR=1.10 +SHADOW_GATE_MIN_DIRECTION_ACCURACY=0.52 +SHADOW_GATE_MAX_BRIER=0.25 DATABASE_PATH=runtime/tradebot.sqlite3 LOG_PATH=runtime/tradebot.log diff --git a/README.md b/README.md index 971bb4f..4d3663f 100644 --- a/README.md +++ b/README.md @@ -247,6 +247,9 @@ Live-исполнение ведет журнал order intent до отправ - `GET /api/status` — статус бота, account snapshot, позиции. - `GET /api/markets` — пары, ticker, свечи, инструменты. - `GET /api/training/market-observations?symbol=BTCUSDT&after_id=0&limit=5000` — защищённая training-token выгрузка L1-наблюдений. +- `GET /api/training/market-observations/manifest` — training-token manifest для инкрементальной синхронизации forward L1-данных. +- `GET /api/training/shadow` — состояние изолированной shadow-модели и повторного forward-gate. +- `POST /api/training/shadow/promote` — атомарное продвижение shadow-модели; возвращает `409`, пока forward-gate не пройден. - `GET /api/trades` — последние сделки. - `GET /api/signals` — последние сигналы стратегии. - `GET /api/events` — события. diff --git a/android/TradeBotMonitor/app/build.gradle.kts b/android/TradeBotMonitor/app/build.gradle.kts index ad6f1e0..55c6713 100644 --- a/android/TradeBotMonitor/app/build.gradle.kts +++ b/android/TradeBotMonitor/app/build.gradle.kts @@ -10,8 +10,8 @@ android { applicationId = "xyz.kusoft.tradebotmonitor" minSdk = 26 targetSdk = 37 - versionCode = 23 - versionName = "0.5.1" + versionCode = 24 + versionName = "0.5.2" } } diff --git a/android/TradeBotMonitor/app/src/main/java/xyz/kusoft/tradebotmonitor/MainActivity.kt b/android/TradeBotMonitor/app/src/main/java/xyz/kusoft/tradebotmonitor/MainActivity.kt index c8b869c..0cfd014 100644 --- a/android/TradeBotMonitor/app/src/main/java/xyz/kusoft/tradebotmonitor/MainActivity.kt +++ b/android/TradeBotMonitor/app/src/main/java/xyz/kusoft/tradebotmonitor/MainActivity.kt @@ -1104,6 +1104,8 @@ class MainActivity : Activity() { addView(trainingComputerPanel(retrain).top(dp(12))) addView(thinDivider().top(dp(12))) addView(trainingProcessPanel(coordination)) + addView(orderbookStagePanel(coordination).top(dp(10))) + addView(shadowStagePanel(retrain.optJSONObject("shadow") ?: JSONObject()).top(dp(10))) if (displayedEvaluation.optJSONObject("candidate") != null) { addView(guardSummaryPanel(displayedEvaluation).top(dp(10))) } @@ -1155,6 +1157,36 @@ class MainActivity : Activity() { } } + private fun orderbookStagePanel(coordination: JSONObject): View = + LinearLayout(this).apply { + orientation = LinearLayout.VERTICAL + val summary = coordination.optJSONObject("latest_job")?.optJSONObject("summary") ?: JSONObject() + val state = summary.optStringClean("state") + val minimum = summary.optInt("minimum_covered_buckets", 240) + val eligible = summary.optInt("eligible_symbol_count", 0) + val requiredSymbols = summary.optInt("minimum_symbols", 2) + val coverage = summary.optJSONObject("covered_buckets_by_symbol") ?: JSONObject() + val bestCoverage = coverage.keys().asSequence().map { coverage.optInt(it, 0) }.maxOrNull() ?: 0 + addView(keyValueLine("Forward-стакан", if (state == "ready") "готов к обучению" else "накапливается", if (state == "ready") palette.green else palette.amber)) + addView(keyValueLine("Лучшее покрытие", "$bestCoverage / $minimum свечей").top(dp(4))) + addView(keyValueLine("Готовые пары", "$eligible / $requiredSymbols").top(dp(4))) + } + + private fun shadowStagePanel(shadow: JSONObject): View = + LinearLayout(this).apply { + orientation = LinearLayout.VERTICAL + val state = shadow.optStringClean("state").ifBlank { "нет модели" } + val color = when (state) { + "passed" -> palette.green + "failed" -> palette.red + else -> palette.amber + } + addView(keyValueLine("Shadow gate", state, color)) + addView(keyValueLine("Forward-прогнозы", "${shadow.optInt("settled_predictions", 0)} / ${shadow.optInt("total_predictions", 0)}").top(dp(4))) + addView(keyValueLine("Shadow P&L", signedPercent(shadow.optDouble("total_net_percent", 0.0)), colorForSigned(shadow.optDouble("total_net_percent", 0.0))).top(dp(4))) + addView(keyValueLine("Direction / Brier", "${percent(shadow.optDouble("direction_accuracy", 0.0) * 100.0, 1)} / ${number(shadow.optDouble("brier", 0.0), 4)}").top(dp(4))) + } + private fun guardSummaryPanel(retrain: JSONObject): View = LinearLayout(this).apply { val accepted = retrain.optBoolean("accepted", false) diff --git a/crypto_spot_bot/__init__.py b/crypto_spot_bot/__init__.py index 8e935c3..8c8e097 100644 --- a/crypto_spot_bot/__init__.py +++ b/crypto_spot_bot/__init__.py @@ -1,3 +1,3 @@ """Crypto spot trading bot package.""" -__version__ = "1.0.3" +__version__ = "1.1.0" diff --git a/crypto_spot_bot/bot.py b/crypto_spot_bot/bot.py index a5d40cb..f6ddd0b 100644 --- a/crypto_spot_bot/bot.py +++ b/crypto_spot_bot/bot.py @@ -2,6 +2,7 @@ from __future__ import annotations import asyncio import logging +import math import sqlite3 from datetime import datetime @@ -14,7 +15,7 @@ from crypto_spot_bot.models import BotStatus, Signal, Ticker, utc_now from crypto_spot_bot.patterns import PatternAnalyzer from crypto_spot_bot.strategy import SpotStrategy, torch_model_readiness_reasons from crypto_spot_bot.storage import Storage -from crypto_spot_bot.time_series import TimeSeriesForecaster +from crypto_spot_bot.time_series import TimeSeriesForecaster, _barrier_outcome logger = logging.getLogger(__name__) @@ -31,6 +32,7 @@ class CryptoSpotBot: pattern_analyzer: PatternAnalyzer, learner: TradeLearner, forecaster: TimeSeriesForecaster | None = None, + shadow_forecaster: TimeSeriesForecaster | None = None, llm_advisor=None, ): self.settings = settings @@ -41,6 +43,7 @@ class CryptoSpotBot: self.pattern_analyzer = pattern_analyzer self.learner = learner self.forecaster = forecaster + self.shadow_forecaster = shadow_forecaster self.llm_advisor = llm_advisor self.running = False self.started_at: datetime | None = None @@ -52,6 +55,8 @@ class CryptoSpotBot: self._last_reconciliation_at: datetime | None = None self._last_prune_at: datetime | None = None self._consecutive_loop_errors = 0 + self._orderbook_feature_cache_key: tuple[tuple[str, int], ...] = () + self._orderbook_feature_cache: dict[str, dict[int, dict[str, float]]] = {} async def start(self) -> None: if self.running: @@ -416,6 +421,25 @@ class CryptoSpotBot: self.market.patterns = patterns def _update_forecasts(self) -> None: + cache_key = tuple( + (symbol, rows[-1].timestamp if rows else 0) + for symbol, rows in sorted(self.market.candles.items()) + ) + earliest_timestamp = min( + (rows[0].timestamp for rows in self.market.candles.values() if rows), + default=0, + ) + if cache_key != self._orderbook_feature_cache_key: + orderbook_features, _manifest = self.storage.recent_aggregated_orderbook_features( + interval=self.settings.base_interval, + symbols=self.market.symbols, + after_timestamp_ms=earliest_timestamp, + min_samples_per_bucket=20, + ) + self._orderbook_feature_cache = orderbook_features + self._orderbook_feature_cache_key = cache_key + else: + orderbook_features = self._orderbook_feature_cache if ( self.forecaster is None or not self.settings.time_series_forecast_enabled @@ -429,8 +453,98 @@ class CryptoSpotBot: symbol=symbol, market_candles=self.market.candles, trend_candles=self.market.trend_candles.get(symbol, []), + orderbook_features=orderbook_features, ).as_dict() self.market.forecasts = forecasts + self._update_shadow_forecasts(orderbook_features) + + def _update_shadow_forecasts( + self, + orderbook_features: dict[str, dict[int, dict[str, float]]], + ) -> None: + if self.shadow_forecaster is None: + self.market.shadow_forecasts = {} + return + model_sha256 = self.shadow_forecaster.artifact_sha256() + if not model_sha256: + self.market.shadow_forecasts = {} + return + forecasts: dict[str, dict] = {} + for symbol in self.market.symbols: + candles = self.market.candles.get(symbol, []) + forecast = self.shadow_forecaster.forecast( + candles, + symbol=symbol, + market_candles=self.market.candles, + trend_candles=self.market.trend_candles.get(symbol, []), + orderbook_features=orderbook_features, + ).as_dict() + forecast["shadow"] = True + forecast["model_sha256"] = model_sha256 + forecasts[symbol] = forecast + self._record_and_settle_shadow(symbol, candles, forecast, model_sha256) + self.market.shadow_forecasts = forecasts + + def _record_and_settle_shadow( + self, + symbol: str, + candles: list, + forecast: dict, + model_sha256: str, + ) -> None: + if candles and forecast.get("usable"): + probability = float( + forecast.get("probability_take_profit_first") + if forecast.get("probability_take_profit_first") is not None + else forecast.get("probability_up", 0.5) + ) + expected = float(forecast.get("expected_return_percent", 0.0) or 0.0) + eligible = bool( + not forecast.get("block_entry") + and expected >= float(forecast.get("calibrated_min_edge_percent", 0.0) or 0.0) + and probability >= float(forecast.get("calibrated_min_probability_up", 0.5) or 0.5) + ) + self.storage.insert_shadow_prediction( + model_sha256=model_sha256, + symbol=symbol, + forecast_timestamp_ms=candles[-1].timestamp, + horizon=max(1, int(forecast.get("horizon", 1) or 1)), + reference_price=float(candles[-1].close), + expected_return_percent=expected, + probability_up=probability, + eligible_signal=eligible, + ) + if not candles: + return + indexes = {candle.timestamp: index for index, candle in enumerate(candles)} + round_trip_cost = 2.0 * ( + float(self.settings.taker_fee_rate) + float(self.settings.slippage_rate) + ) + for row in self.storage.pending_shadow_predictions( + model_sha256=model_sha256, + symbol=symbol, + ): + index = indexes.get(int(row.get("forecast_timestamp_ms", 0) or 0)) + horizon = max(1, int(row.get("horizon", 1) or 1)) + if index is None or index + horizon >= len(candles): + continue + outcome = _barrier_outcome( + candles, + end_index=index, + horizon=horizon, + stop_loss_percent=float(self.settings.stop_loss_percent), + take_profit_percent=float(self.settings.take_profit_percent), + round_trip_cost=round_trip_cost, + ) + if outcome is None: + continue + actual_log_return, take_profit_first = outcome + actual_return_percent = (math.exp(actual_log_return) - 1.0) * 100.0 + self.storage.settle_shadow_prediction( + int(row["id"]), + actual_return_percent=actual_return_percent, + take_profit_first=take_profit_first >= 0.5, + ) def status(self) -> BotStatus: live_ready = self.settings.live_ready diff --git a/crypto_spot_bot/dashboard.py b/crypto_spot_bot/dashboard.py index d423c17..a4ce191 100644 --- a/crypto_spot_bot/dashboard.py +++ b/crypto_spot_bot/dashboard.py @@ -20,6 +20,7 @@ from crypto_spot_bot.learning import TradeLearner from crypto_spot_bot.market_data import MarketData from crypto_spot_bot.patterns import PatternAnalyzer from crypto_spot_bot.reconciliation import reconciliation_snapshot +from crypto_spot_bot.shadow import shadow_gate_snapshot from crypto_spot_bot.storage import Storage from crypto_spot_bot.strategy import SpotStrategy from crypto_spot_bot.time_series import TimeSeriesForecaster @@ -47,7 +48,23 @@ def create_app(settings: Settings | None = None) -> FastAPI: pattern_analyzer = PatternAnalyzer() learner = TradeLearner(settings, storage) forecaster = TimeSeriesForecaster(settings) - bot = CryptoSpotBot(settings, storage, market, broker, strategy, pattern_analyzer, learner, forecaster) + runtime_dir = settings.time_series_lstm_model_path.parent + shadow_forecaster = TimeSeriesForecaster( + settings, + model_path=runtime_dir / "lstm_forecaster.shadow.json", + calibration_path=runtime_dir / "torch_shadow_calibration.json", + ) + bot = CryptoSpotBot( + settings, + storage, + market, + broker, + strategy, + pattern_analyzer, + learner, + forecaster, + shadow_forecaster, + ) training = TrainingCoordinator(settings.time_series_lstm_model_path.parent) authorizer = ApiAuthorizer(settings) @@ -143,12 +160,31 @@ def create_app(settings: Settings | None = None) -> FastAPI: async def retrain(_: None = Depends(authorizer.require)) -> dict[str, Any]: data = _runtime_json(settings, "torch_retrain_guard.json") data["coordination"] = training.status() + data["shadow"] = shadow_gate_snapshot(storage, shadow_forecaster.artifact_sha256()) return data @app.get("/api/training/status") async def training_status(_: None = Depends(authorizer.require)) -> dict[str, Any]: return training.status() + @app.get("/api/training/shadow") + async def training_shadow_status( + _: None = Depends(authorizer.require), + ) -> dict[str, Any]: + return shadow_gate_snapshot(storage, shadow_forecaster.artifact_sha256()) + + @app.post("/api/training/shadow/promote") + async def training_shadow_promote( + _: None = Depends(authorizer.require), + ) -> dict[str, Any]: + gate = shadow_gate_snapshot(storage, shadow_forecaster.artifact_sha256()) + if not gate.get("passed"): + raise HTTPException(status_code=409, detail={"message": "shadow forward gate has not passed", "gate": gate}) + try: + return training.promote_shadow(gate) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + @app.get("/api/training/market-observations") async def training_market_observations( symbol: str, @@ -170,6 +206,16 @@ def create_app(settings: Settings | None = None) -> FastAPI: "next_after_id": int(items[-1]["id"]) if items else max(0, after_id), } + @app.get("/api/training/market-observations/manifest") + async def training_market_observation_manifest( + _: None = Depends(authorizer.require_training), + ) -> dict[str, Any]: + items = storage.market_observation_manifest() + return { + "items": items, + "total_samples": sum(int(item.get("samples", 0) or 0) for item in items), + } + @app.post("/api/training/retrain") async def training_retrain( payload: dict[str, Any] | None = None, @@ -233,6 +279,10 @@ def create_app(settings: Settings | None = None) -> FastAPI: row_limit = 220 retrain_data = _runtime_json(settings, "torch_retrain_guard.json") retrain_data["coordination"] = training.status() + retrain_data["shadow"] = shadow_gate_snapshot( + storage, + shadow_forecaster.artifact_sha256(), + ) return { "health": { "ok": True, diff --git a/crypto_spot_bot/market_data.py b/crypto_spot_bot/market_data.py index bee5d33..38951f4 100644 --- a/crypto_spot_bot/market_data.py +++ b/crypto_spot_bot/market_data.py @@ -55,6 +55,7 @@ class MarketData: self.orderbook_metrics: dict[str, dict[str, Any]] = {} self.patterns: dict[str, dict[str, Any]] = {} self.forecasts: dict[str, dict[str, Any]] = {} + self.shadow_forecasts: dict[str, dict[str, Any]] = {} self.last_rest_refresh_at: datetime | None = None self.last_ws_message_at: datetime | None = None self.ws_connected = False @@ -393,6 +394,7 @@ class MarketData: "trend_candles": [candle.as_dict() for candle in self.trend_candles.get(symbol, [])[-5:]], "pattern": self.patterns.get(symbol), "forecast": self.forecasts.get(symbol), + "shadow_forecast": self.shadow_forecasts.get(symbol), "orderbook": self.orderbook_metrics.get(symbol), "quality": analyze_symbol_quality( symbol=symbol, diff --git a/crypto_spot_bot/orderbook_features.py b/crypto_spot_bot/orderbook_features.py new file mode 100644 index 0000000..893899b --- /dev/null +++ b/crypto_spot_bot/orderbook_features.py @@ -0,0 +1,180 @@ +from __future__ import annotations + +import math +import sqlite3 +from collections import defaultdict +from datetime import datetime +from pathlib import Path +from typing import Any, Iterable + + +ORDERBOOK_FEATURES = ( + "l1_imbalance_mean", + "l1_imbalance_std", + "l1_spread_bps_mean", + "l1_spread_bps_p90", + "l1_microprice_deviation_bps_mean", + "l1_microprice_deviation_bps_std", + "l1_sample_count_log1p", +) + + +def interval_milliseconds(interval: str) -> int: + normalized = str(interval).strip().upper() + if normalized.isdigit(): + return max(1, int(normalized)) * 60_000 + units = { + "D": 86_400_000, + "W": 7 * 86_400_000, + "M": 30 * 86_400_000, + } + return units.get(normalized, 0) + + +def load_orderbook_feature_map( + path: str | Path, + *, + interval: str, + symbols: Iterable[str] | None = None, + min_samples_per_bucket: int = 20, +) -> tuple[dict[str, dict[int, dict[str, float]]], dict[str, dict[str, Any]]]: + database_path = Path(path) + if not database_path.is_file(): + return {}, {} + selected = sorted({str(symbol).strip().upper() for symbol in symbols or [] if str(symbol).strip()}) + query = ( + "SELECT symbol, bid_price, bid_size, ask_price, ask_size, mid_price, " + "microprice, spread_bps, imbalance, source_timestamp_ms, created_at " + "FROM market_observations" + ) + parameters: list[Any] = [] + if selected: + placeholders = ",".join("?" for _ in selected) + query += f" WHERE symbol IN ({placeholders})" + parameters.extend(selected) + query += " ORDER BY symbol, source_timestamp_ms, created_at" + with sqlite3.connect(database_path) as connection: + connection.row_factory = sqlite3.Row + try: + rows = connection.execute(query, parameters).fetchall() + except sqlite3.Error: + return {}, {} + return aggregate_orderbook_observations( + (dict(row) for row in rows), + interval=interval, + min_samples_per_bucket=min_samples_per_bucket, + ) + + +def aggregate_orderbook_observations( + rows: Iterable[dict[str, Any]], + *, + interval: str, + min_samples_per_bucket: int = 20, +) -> tuple[dict[str, dict[int, dict[str, float]]], dict[str, dict[str, Any]]]: + interval_ms = interval_milliseconds(interval) + if interval_ms <= 0: + raise ValueError(f"unsupported orderbook aggregation interval: {interval}") + minimum = max(1, int(min_samples_per_bucket)) + buckets: dict[tuple[str, int], list[tuple[float, float, float]]] = defaultdict(list) + raw_counts: dict[str, int] = defaultdict(int) + first_timestamp: dict[str, int] = {} + last_timestamp: dict[str, int] = {} + for row in rows: + symbol = str(row.get("symbol") or "").strip().upper() + timestamp_ms = _observation_timestamp_ms(row) + mid_price = _float(row.get("mid_price")) + microprice = _float(row.get("microprice"), mid_price) + spread_bps = max(0.0, _float(row.get("spread_bps"))) + imbalance = max(-1.0, min(1.0, _float(row.get("imbalance")))) + if not symbol or timestamp_ms <= 0 or mid_price <= 0: + continue + microprice_deviation_bps = ((microprice - mid_price) / mid_price) * 10_000.0 + if not all(math.isfinite(value) for value in (imbalance, spread_bps, microprice_deviation_bps)): + continue + bucket_timestamp = (timestamp_ms // interval_ms) * interval_ms + buckets[(symbol, bucket_timestamp)].append( + (imbalance, spread_bps, microprice_deviation_bps) + ) + raw_counts[symbol] += 1 + first_timestamp[symbol] = min(first_timestamp.get(symbol, timestamp_ms), timestamp_ms) + last_timestamp[symbol] = max(last_timestamp.get(symbol, timestamp_ms), timestamp_ms) + + features: dict[str, dict[int, dict[str, float]]] = defaultdict(dict) + rejected_buckets: dict[str, int] = defaultdict(int) + for (symbol, bucket_timestamp), samples in sorted(buckets.items()): + if len(samples) < minimum: + rejected_buckets[symbol] += 1 + continue + imbalances = [sample[0] for sample in samples] + spreads = [sample[1] for sample in samples] + microprice_deviations = [sample[2] for sample in samples] + features[symbol][bucket_timestamp] = { + "l1_imbalance_mean": _mean(imbalances), + "l1_imbalance_std": _standard_deviation(imbalances), + "l1_spread_bps_mean": _mean(spreads), + "l1_spread_bps_p90": _percentile(spreads, 0.90), + "l1_microprice_deviation_bps_mean": _mean(microprice_deviations), + "l1_microprice_deviation_bps_std": _standard_deviation(microprice_deviations), + "l1_sample_count_log1p": math.log1p(len(samples)), + } + + manifest: dict[str, dict[str, Any]] = {} + all_symbols = sorted(set(raw_counts) | set(features)) + for symbol in all_symbols: + accepted = features.get(symbol, {}) + manifest[symbol] = { + "raw_samples": raw_counts.get(symbol, 0), + "covered_buckets": len(accepted), + "rejected_buckets": rejected_buckets.get(symbol, 0), + "first_timestamp_ms": first_timestamp.get(symbol, 0), + "last_timestamp_ms": last_timestamp.get(symbol, 0), + "min_samples_per_bucket": minimum, + } + return {symbol: dict(rows) for symbol, rows in features.items()}, manifest + + +def _observation_timestamp_ms(row: dict[str, Any]) -> int: + source_timestamp = int(_float(row.get("source_timestamp_ms"))) + if source_timestamp > 0: + return source_timestamp + raw = str(row.get("created_at") or "").strip() + if not raw: + return 0 + try: + parsed = datetime.fromisoformat(raw.replace("Z", "+00:00")) + except ValueError: + return 0 + return int(parsed.timestamp() * 1000) + + +def _mean(values: list[float]) -> float: + return sum(values) / len(values) if values else 0.0 + + +def _standard_deviation(values: list[float]) -> float: + if len(values) < 2: + return 0.0 + mean = _mean(values) + return math.sqrt(sum((value - mean) ** 2 for value in values) / len(values)) + + +def _percentile(values: list[float], quantile: float) -> float: + if not values: + return 0.0 + ordered = sorted(values) + position = max(0.0, min(1.0, quantile)) * (len(ordered) - 1) + lower = int(math.floor(position)) + upper = int(math.ceil(position)) + if lower == upper: + return ordered[lower] + fraction = position - lower + return ordered[lower] * (1.0 - fraction) + ordered[upper] * fraction + + +def _float(value: Any, default: float = 0.0) -> float: + try: + result = float(value) + except (TypeError, ValueError): + return default + return result if math.isfinite(result) else default diff --git a/crypto_spot_bot/shadow.py b/crypto_spot_bot/shadow.py new file mode 100644 index 0000000..4816334 --- /dev/null +++ b/crypto_spot_bot/shadow.py @@ -0,0 +1,94 @@ +from __future__ import annotations + +import math +import os +from typing import Any + +from crypto_spot_bot.storage import Storage + + +def shadow_gate_snapshot(storage: Storage, model_sha256: str) -> dict[str, Any]: + minimum_settled = _int_env("SHADOW_GATE_MIN_SETTLED", 300) + minimum_eligible = _int_env("SHADOW_GATE_MIN_ELIGIBLE", 30) + minimum_symbols = _int_env("SHADOW_GATE_MIN_SYMBOLS", 2) + minimum_profit_factor = _float_env("SHADOW_GATE_MIN_PROFIT_FACTOR", 1.10) + minimum_direction_accuracy = _float_env("SHADOW_GATE_MIN_DIRECTION_ACCURACY", 0.52) + maximum_brier = _float_env("SHADOW_GATE_MAX_BRIER", 0.25) + rows = storage.shadow_prediction_rows(model_sha256=model_sha256) if model_sha256 else [] + settled = [row for row in rows if row.get("settled_at")] + eligible = [row for row in settled if bool(row.get("eligible_signal"))] + eligible_returns = [float(row.get("actual_return_percent", 0.0) or 0.0) for row in eligible] + gross_profit = sum(max(0.0, value) for value in eligible_returns) + gross_loss = abs(sum(min(0.0, value) for value in eligible_returns)) + profit_factor = gross_profit / gross_loss if gross_loss > 1e-12 else (float("inf") if gross_profit > 0 else 0.0) + correct = sum( + 1 + for row in settled + if (float(row.get("expected_return_percent", 0.0) or 0.0) >= 0) + == (float(row.get("actual_return_percent", 0.0) or 0.0) >= 0) + ) + direction_accuracy = correct / len(settled) if settled else 0.0 + brier_values = [ + ( + max(0.0, min(1.0, float(row.get("probability_up", 0.5) or 0.5))) + - float(int(row.get("take_profit_first", 0) or 0)) + ) + ** 2 + for row in settled + if row.get("take_profit_first") is not None + ] + brier = sum(brier_values) / len(brier_values) if brier_values else 1.0 + symbols = sorted({str(row.get("symbol") or "") for row in eligible if row.get("symbol")}) + checks = { + "minimum_settled": len(settled) >= minimum_settled, + "minimum_eligible": len(eligible) >= minimum_eligible, + "minimum_symbols": len(symbols) >= minimum_symbols, + "positive_average_net": bool(eligible_returns) and sum(eligible_returns) / len(eligible_returns) > 0.0, + "profit_factor": profit_factor >= minimum_profit_factor, + "direction_accuracy": direction_accuracy >= minimum_direction_accuracy, + "brier": brier <= maximum_brier, + } + enough_data = checks["minimum_settled"] and checks["minimum_eligible"] and checks["minimum_symbols"] + passed = enough_data and all(checks.values()) + state = "passed" if passed else ("failed" if enough_data else "collecting") + return { + "available": bool(model_sha256), + "model_sha256": model_sha256, + "state": state, + "passed": passed, + "active_model_unchanged": True, + "total_predictions": len(rows), + "pending_predictions": len(rows) - len(settled), + "settled_predictions": len(settled), + "eligible_predictions": len(eligible), + "eligible_symbols": symbols, + "average_net_percent": round(sum(eligible_returns) / len(eligible_returns), 6) if eligible_returns else 0.0, + "total_net_percent": round(sum(eligible_returns), 6), + "win_rate": round(sum(value > 0 for value in eligible_returns) / len(eligible_returns), 6) if eligible_returns else 0.0, + "profit_factor": round(profit_factor, 6) if math.isfinite(profit_factor) else None, + "direction_accuracy": round(direction_accuracy, 6), + "brier": round(brier, 6), + "criteria": { + "minimum_settled": minimum_settled, + "minimum_eligible": minimum_eligible, + "minimum_symbols": minimum_symbols, + "minimum_profit_factor": minimum_profit_factor, + "minimum_direction_accuracy": minimum_direction_accuracy, + "maximum_brier": maximum_brier, + }, + "checks": checks, + } + + +def _int_env(name: str, default: int) -> int: + try: + return max(1, int(os.environ.get(name, str(default)))) + except ValueError: + return default + + +def _float_env(name: str, default: float) -> float: + try: + return float(os.environ.get(name, str(default))) + except ValueError: + return default diff --git a/crypto_spot_bot/storage.py b/crypto_spot_bot/storage.py index dc5766b..a07885c 100644 --- a/crypto_spot_bot/storage.py +++ b/crypto_spot_bot/storage.py @@ -9,6 +9,7 @@ from pathlib import Path from typing import Any, Iterator from crypto_spot_bot.models import Position, Signal, Trade, utc_now +from crypto_spot_bot.orderbook_features import aggregate_orderbook_observations, load_orderbook_feature_map MAX_SIGNAL_DIAGNOSTICS_BYTES = 4 * 1024 @@ -19,6 +20,7 @@ MAX_RUNTIME_ROWS = { "events": 20_000, "llm_advice": 20_000, "market_observations": 1_200_000, + "shadow_predictions": 250_000, } _STORED_FORECAST_KEYS = { "enabled", @@ -187,6 +189,22 @@ class Storage: source_timestamp_ms INTEGER NOT NULL DEFAULT 0, created_at TEXT NOT NULL ); + CREATE TABLE IF NOT EXISTS shadow_predictions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + model_sha256 TEXT NOT NULL, + symbol TEXT NOT NULL, + forecast_timestamp_ms INTEGER NOT NULL, + horizon INTEGER NOT NULL, + reference_price REAL NOT NULL, + expected_return_percent REAL NOT NULL, + probability_up REAL NOT NULL, + eligible_signal INTEGER NOT NULL DEFAULT 0, + created_at TEXT NOT NULL, + settled_at TEXT, + actual_return_percent REAL, + take_profit_first INTEGER, + UNIQUE(model_sha256, symbol, forecast_timestamp_ms, horizon) + ); CREATE INDEX IF NOT EXISTS idx_positions_status_opened ON positions(status, opened_at); CREATE INDEX IF NOT EXISTS idx_trades_closed @@ -203,6 +221,10 @@ class Storage: ON market_observations(symbol, id); CREATE INDEX IF NOT EXISTS idx_market_observations_created ON market_observations(created_at); + CREATE INDEX IF NOT EXISTS idx_market_observations_symbol_source_timestamp + ON market_observations(symbol, source_timestamp_ms); + CREATE INDEX IF NOT EXISTS idx_shadow_predictions_model_status + ON shadow_predictions(model_sha256, settled_at, symbol); """ ) columns = { @@ -542,6 +564,159 @@ class Storage: ).fetchall() return [dict(row) for row in rows] + def market_observation_manifest(self) -> list[dict[str, Any]]: + with self.connect() as conn: + rows = conn.execute( + """ + SELECT symbol, COUNT(*) AS samples, MIN(id) AS min_id, MAX(id) AS max_id, + MIN(source_timestamp_ms) AS first_source_timestamp_ms, + MAX(source_timestamp_ms) AS last_source_timestamp_ms, + MIN(created_at) AS first_created_at, + MAX(created_at) AS last_created_at + FROM market_observations + GROUP BY symbol + ORDER BY symbol + """ + ).fetchall() + return [dict(row) for row in rows] + + def aggregated_orderbook_features( + self, + *, + interval: str, + symbols: list[str] | None = None, + min_samples_per_bucket: int = 20, + ) -> tuple[dict[str, dict[int, dict[str, float]]], dict[str, dict[str, Any]]]: + return load_orderbook_feature_map( + self.path, + interval=interval, + symbols=symbols, + min_samples_per_bucket=min_samples_per_bucket, + ) + + def recent_aggregated_orderbook_features( + self, + *, + interval: str, + symbols: list[str], + after_timestamp_ms: int, + min_samples_per_bucket: int = 20, + ) -> tuple[dict[str, dict[int, dict[str, float]]], dict[str, dict[str, Any]]]: + selected = sorted({symbol.strip().upper() for symbol in symbols if symbol.strip()}) + if not selected: + return {}, {} + placeholders = ",".join("?" for _ in selected) + with self.connect() as conn: + rows = conn.execute( + f""" + SELECT symbol, bid_price, bid_size, ask_price, ask_size, mid_price, + microprice, spread_bps, imbalance, source_timestamp_ms, created_at + FROM market_observations + WHERE symbol IN ({placeholders}) AND source_timestamp_ms >= ? + ORDER BY symbol, source_timestamp_ms + """, + (*selected, max(0, int(after_timestamp_ms))), + ).fetchall() + return aggregate_orderbook_observations( + (dict(row) for row in rows), + interval=interval, + min_samples_per_bucket=min_samples_per_bucket, + ) + + def insert_shadow_prediction( + self, + *, + model_sha256: str, + symbol: str, + forecast_timestamp_ms: int, + horizon: int, + reference_price: float, + expected_return_percent: float, + probability_up: float, + eligible_signal: bool, + ) -> bool: + with self.connect() as conn: + cursor = conn.execute( + """ + INSERT OR IGNORE INTO shadow_predictions ( + model_sha256, symbol, forecast_timestamp_ms, horizon, + reference_price, expected_return_percent, probability_up, + eligible_signal, created_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + model_sha256, + symbol.upper(), + max(0, int(forecast_timestamp_ms)), + max(1, int(horizon)), + max(0.0, float(reference_price)), + float(expected_return_percent), + max(0.0, min(1.0, float(probability_up))), + 1 if eligible_signal else 0, + utc_now().isoformat(), + ), + ) + return bool(cursor.rowcount) + + def pending_shadow_predictions( + self, + *, + model_sha256: str, + symbol: str, + limit: int = 500, + ) -> list[dict[str, Any]]: + with self.connect() as conn: + rows = conn.execute( + """ + SELECT * FROM shadow_predictions + WHERE model_sha256 = ? AND symbol = ? AND settled_at IS NULL + ORDER BY forecast_timestamp_ms + LIMIT ? + """, + (model_sha256, symbol.upper(), max(1, min(5000, int(limit)))), + ).fetchall() + return [dict(row) for row in rows] + + def settle_shadow_prediction( + self, + prediction_id: int, + *, + actual_return_percent: float, + take_profit_first: bool, + ) -> bool: + with self.connect() as conn: + cursor = conn.execute( + """ + UPDATE shadow_predictions + SET settled_at = ?, actual_return_percent = ?, take_profit_first = ? + WHERE id = ? AND settled_at IS NULL + """, + ( + utc_now().isoformat(), + float(actual_return_percent), + 1 if take_profit_first else 0, + int(prediction_id), + ), + ) + return bool(cursor.rowcount) + + def shadow_prediction_rows( + self, + *, + model_sha256: str, + settled_only: bool = False, + limit: int = 250_000, + ) -> list[dict[str, Any]]: + where = "WHERE model_sha256 = ?" + if settled_only: + where += " AND settled_at IS NOT NULL" + with self.connect() as conn: + rows = conn.execute( + f"SELECT * FROM shadow_predictions {where} ORDER BY id DESC LIMIT ?", + (model_sha256, max(1, min(250_000, int(limit)))), + ).fetchall() + return [dict(row) for row in rows] + def insert_equity( self, equity: float, @@ -719,7 +894,14 @@ class Storage: return {} cutoff = (utc_now() - timedelta(days=retention_days)).isoformat() deleted: dict[str, int] = {} - for table in ("signals", "equity", "events", "llm_advice", "market_observations"): + for table in ( + "signals", + "equity", + "events", + "llm_advice", + "market_observations", + "shadow_predictions", + ): with self.connect() as conn: max_id_row = conn.execute(f"SELECT MAX(id) AS value FROM {table}").fetchone() max_id = int(max_id_row["value"] or 0) if max_id_row else 0 diff --git a/crypto_spot_bot/time_series.py b/crypto_spot_bot/time_series.py index 9a3bcdb..eeaf95b 100644 --- a/crypto_spot_bot/time_series.py +++ b/crypto_spot_bot/time_series.py @@ -2,13 +2,16 @@ from __future__ import annotations import json import math +import hashlib from bisect import bisect_right from dataclasses import asdict, dataclass, field from datetime import UTC, datetime +from pathlib import Path from typing import Any from crypto_spot_bot.config import Settings from crypto_spot_bot.models import Candle +from crypto_spot_bot.orderbook_features import ORDERBOOK_FEATURES DEFAULT_TORCH_FEATURES = ( @@ -170,8 +173,18 @@ class TimeSeriesForecast: class TimeSeriesForecaster: - def __init__(self, settings: Settings): + def __init__( + self, + settings: Settings, + *, + model_path: Path | None = None, + calibration_path: Path | None = None, + ): self.settings = settings + self.model_path = model_path or settings.time_series_lstm_model_path + self.calibration_path = calibration_path or ( + self.model_path.parent / "torch_threshold_calibration.json" + ) self._lstm_artifact_mtime: float | None = None self._lstm_artifact: dict[str, Any] = {} self._calibration_mtime: float | None = None @@ -184,6 +197,7 @@ class TimeSeriesForecaster: *, market_candles: dict[str, list[Candle]] | None = None, trend_candles: list[Candle] | None = None, + orderbook_features: dict[str, dict[int, dict[str, float]]] | None = None, ) -> TimeSeriesForecast: if not self.settings.time_series_forecast_enabled: return _empty_forecast(False, "time-series forecast is disabled") @@ -225,6 +239,7 @@ class TimeSeriesForecaster: symbol=symbol, market_candles=market_candles, trend_candles=trend_candles, + orderbook_features=orderbook_features, ) if entry else [] @@ -413,7 +428,7 @@ class TimeSeriesForecaster: def _load_lstm_artifact(self) -> dict[str, Any]: if not self.settings.time_series_lstm_enabled: return {} - path = self.settings.time_series_lstm_model_path + path = self.model_path try: stat = path.stat() except OSError: @@ -431,7 +446,7 @@ class TimeSeriesForecaster: return self._lstm_artifact def _load_quality_gate(self) -> dict[str, Any]: - path = self.settings.time_series_lstm_model_path.parent / "torch_threshold_calibration.json" + path = self.calibration_path try: stat = path.stat() except OSError: @@ -448,6 +463,12 @@ class TimeSeriesForecaster: self._calibration_mtime = stat.st_mtime return self._quality_gate + def artifact_sha256(self) -> str: + try: + return hashlib.sha256(self.model_path.read_bytes()).hexdigest() + except OSError: + return "" + def _empty_forecast(enabled: bool, reason: str) -> TimeSeriesForecast: return TimeSeriesForecast( @@ -554,6 +575,7 @@ def _feature_matrix( symbol: str | None = None, market_candles: dict[str, list[Candle]] | None = None, trend_candles: list[Candle] | None = None, + orderbook_features: dict[str, dict[int, dict[str, float]]] | None = None, ) -> list[list[float]]: names = list(feature_names or DEFAULT_TORCH_FEATURES) context = _feature_context( @@ -561,6 +583,7 @@ def _feature_matrix( symbol=symbol, market_candles=market_candles, trend_candles=trend_candles, + orderbook_features=orderbook_features, ) rows: list[list[float]] = [] for index, candle in enumerate(candles): @@ -574,6 +597,7 @@ def _feature_context( symbol: str | None, market_candles: dict[str, list[Candle]] | None, trend_candles: list[Candle] | None, + orderbook_features: dict[str, dict[int, dict[str, float]]] | None, ) -> dict[str, Any]: market_candles = market_candles or {} normalized_market = {key.upper(): value for key, value in market_candles.items()} @@ -595,6 +619,9 @@ def _feature_context( "context_indexes": context_indexes, "trend_candles": trend_rows, "trend_positions": trend_positions, + "orderbook_features": { + key.upper(): value for key, value in (orderbook_features or {}).items() + }, } @@ -603,6 +630,12 @@ def _feature_value(name: str, candles: list[Candle], index: int, candle: Candle, previous = candles[index - 1] if index >= 1 else candle if name.startswith("symbol_is_"): return 1.0 if context.get("symbol") == name.removeprefix("symbol_is_").upper() else 0.0 + if name in ORDERBOOK_FEATURES: + symbol_features = (context.get("orderbook_features") or {}).get( + context.get("symbol"), {} + ) + values = symbol_features.get(candle.timestamp, {}) + return _safe_feature(float(values.get(name, 0.0) or 0.0)) if name == "return_1": return _log_change(candle.close, previous.close) if name == "return_3": diff --git a/crypto_spot_bot/training_coordination.py b/crypto_spot_bot/training_coordination.py index fa2e810..0d5ac51 100644 --- a/crypto_spot_bot/training_coordination.py +++ b/crypto_spot_bot/training_coordination.py @@ -17,11 +17,17 @@ from threading import Lock from typing import Any -ALLOWED_TRAINING_ARTIFACTS = { +ACTIVE_TRAINING_ARTIFACTS = { "lstm_forecaster.json", "torch_retrain_guard.json", "torch_threshold_calibration.json", } +SHADOW_TRAINING_ARTIFACTS = { + "lstm_forecaster.shadow.json", + "torch_shadow_guard.json", + "torch_shadow_calibration.json", +} +ALLOWED_TRAINING_ARTIFACTS = ACTIVE_TRAINING_ARTIFACTS | SHADOW_TRAINING_ARTIFACTS RUNNING_LEASE_TIMEOUT = timedelta(minutes=10) ONLINE_WINDOW = timedelta(minutes=3) MAX_JOB_ATTEMPTS = 3 @@ -29,7 +35,8 @@ MAX_ARTIFACT_CHUNK_BYTES = 1024 * 1024 # Keep uploads bounded while leaving room for explicitly requested per-symbol bundles. MAX_ARTIFACT_BYTES = 256 * 1024 * 1024 MAX_ARTIFACT_CHUNKS = 1024 -REQUIRED_MODEL_BUNDLE = set(ALLOWED_TRAINING_ARTIFACTS) +REQUIRED_MODEL_BUNDLE = set(ACTIVE_TRAINING_ARTIFACTS) +REQUIRED_SHADOW_BUNDLE = set(SHADOW_TRAINING_ARTIFACTS) class TrainingCoordinator: @@ -46,6 +53,61 @@ class TrainingCoordinator: self._save_state(state) return self._public_status(state) + def promote_shadow(self, forward_gate: dict[str, Any]) -> dict[str, Any]: + with self._lock: + if not bool(forward_gate.get("passed")): + raise ValueError("shadow forward gate has not passed") + shadow_model = self.runtime_dir / "lstm_forecaster.shadow.json" + shadow_calibration = self.runtime_dir / "torch_shadow_calibration.json" + shadow_guard = self.runtime_dir / "torch_shadow_guard.json" + missing = [ + path.name + for path in (shadow_model, shadow_calibration, shadow_guard) + if not path.is_file() + ] + if missing: + raise ValueError("shadow bundle is incomplete: " + ", ".join(missing)) + model_sha256 = hashlib.sha256(shadow_model.read_bytes()).hexdigest() + if str(forward_gate.get("model_sha256") or "") != model_sha256: + raise ValueError("shadow forward gate is bound to another model") + calibration = _read_json(shadow_calibration) + guard = _read_json(shadow_guard) + if calibration.get("artifact_sha256") != model_sha256: + raise ValueError("shadow calibration is not bound to the model") + if not bool(guard.get("accepted")) or guard.get("candidate_artifact_sha256") != model_sha256: + raise ValueError("shadow offline guard is invalid") + + promotion_id = str(uuid.uuid4()) + backup_dir = self.runtime_dir / ".model_backups" / f"{_compact_now()}-shadow-{promotion_id}" + backup_dir.mkdir(parents=True, exist_ok=True) + targets = { + "lstm_forecaster.json": shadow_model, + "torch_threshold_calibration.json": shadow_calibration, + "torch_retrain_guard.json": shadow_guard, + } + for target_name in targets: + current = self.runtime_dir / target_name + if current.is_file(): + shutil.copy2(current, backup_dir / target_name) + for target_name, source in targets.items(): + target_tmp = self.runtime_dir / f".{target_name}.{promotion_id}.promote" + shutil.copy2(source, target_tmp) + os.replace(target_tmp, self.runtime_dir / target_name) + gate_path = self.runtime_dir / "torch_shadow_forward_gate.json" + gate_tmp = gate_path.with_suffix(".tmp") + gate_tmp.write_text( + json.dumps(forward_gate, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + os.replace(gate_tmp, gate_path) + return { + "promoted": True, + "model_sha256": model_sha256, + "promotion_id": promotion_id, + "backup_dir": str(backup_dir), + "promoted_at": _now(), + } + def request_retrain(self, payload: dict[str, Any] | None = None) -> dict[str, Any]: payload = payload or {} with self._lock: @@ -238,8 +300,17 @@ class TrainingCoordinator: self._require_lease(job, payload) success = bool(payload.get("success", payload.get("status") == "completed")) if success and job.get("artifacts"): - promoted = self._validate_and_promote(job_id, job) - job["promoted_artifacts"] = promoted + artifact_names = { + str(item.get("name")) + for item in job.get("artifacts", []) + if isinstance(item, dict) + } + if artifact_names & REQUIRED_SHADOW_BUNDLE: + staged = self._validate_and_stage_shadow(job_id, job) + job["shadow_artifacts"] = staged + else: + promoted = self._validate_and_promote(job_id, job) + job["promoted_artifacts"] = promoted job["status"] = "completed" if success else "failed" job["phase"] = "completed" if success else "failed" job["progress_percent"] = 100 if success else _coerce_percent(payload.get("progress_percent"), job.get("progress_percent", 0)) @@ -247,7 +318,9 @@ class TrainingCoordinator: job["message"] = str(payload.get("message") or "") if isinstance(payload.get("summary"), dict): job["summary"] = payload["summary"] - if isinstance(payload["summary"].get("accepted"), bool): + if str(payload["summary"].get("state") or "").startswith("collecting"): + job["model_decision"] = "collecting" + elif isinstance(payload["summary"].get("accepted"), bool): job["model_decision"] = ( "accepted" if payload["summary"]["accepted"] else "rejected" ) @@ -318,6 +391,60 @@ class TrainingCoordinator: _remove_tree(self.upload_root / job_id) return promoted + def _validate_and_stage_shadow(self, job_id: str, job: dict[str, Any]) -> list[dict[str, Any]]: + ready_dir = self.upload_root / job_id / "ready" + staged = {path.name for path in ready_dir.iterdir() if path.is_file()} if ready_dir.is_dir() else set() + missing = REQUIRED_SHADOW_BUNDLE - staged + if missing: + raise ValueError("shadow training bundle is incomplete: " + ", ".join(sorted(missing))) + + model_path = ready_dir / "lstm_forecaster.shadow.json" + calibration_path = ready_dir / "torch_shadow_calibration.json" + guard_path = ready_dir / "torch_shadow_guard.json" + model = _read_json(model_path) + calibration = _read_json(calibration_path) + guard = _read_json(guard_path) + if model.get("type") != "pytorch_recurrent_forecaster": + raise ValueError("shadow candidate model type is invalid") + symbols = model.get("symbols") + if not isinstance(symbols, dict) or not symbols: + raise ValueError("shadow candidate model has no symbol models") + _validate_symbol_models(symbols) + model_sha256 = hashlib.sha256(model_path.read_bytes()).hexdigest() + if calibration.get("artifact_sha256") != model_sha256: + raise ValueError("shadow calibration is not bound to the uploaded model") + if not bool(guard.get("accepted")): + raise ValueError("shadow candidate did not pass the offline guard") + if guard.get("candidate_artifact_sha256") != model_sha256: + raise ValueError("shadow guard is not bound to the uploaded model") + validation = calibration.get("validation") + if not isinstance(validation, dict) or not _validation_passed(validation): + raise ValueError("shadow candidate offline quality gate did not pass") + if validation.get("protocol") != "untouched_model_holdout_with_threshold_walk_forward": + raise ValueError("shadow candidate validation protocol is not an untouched holdout") + + self.runtime_dir.mkdir(parents=True, exist_ok=True) + artifact_rows = { + str(item.get("name")): item + for item in job.get("artifacts", []) + if isinstance(item, dict) + } + installed: list[dict[str, Any]] = [] + for name in sorted(REQUIRED_SHADOW_BUNDLE): + target_tmp = self.runtime_dir / f".{name}.{job_id}.stage" + shutil.copy2(ready_dir / name, target_tmp) + os.replace(target_tmp, self.runtime_dir / name) + row = artifact_rows.get(name, {}) + installed.append( + { + "name": name, + "sha256": row.get("sha256", ""), + "staged_at": _now(), + } + ) + _remove_tree(self.upload_root / job_id) + return installed + def _load_state(self) -> dict[str, Any]: try: data = json.loads(self.state_path.read_text(encoding="utf-8")) @@ -464,6 +591,10 @@ def _safe_parameters(value: Any) -> dict[str, Any]: "interval", "pooled", "resume_candidate", + "use_orderbook", + "orderbook_min_samples_per_bucket", + "orderbook_min_covered_buckets", + "orderbook_min_symbols", } result = {key: value[key] for key in allowed if key in value} for key, low, high in ( @@ -475,6 +606,9 @@ def _safe_parameters(value: Any) -> dict[str, Any]: ("horizon", 1, 96), ("patience", 1, 50), ("seed", 1, 2_147_483_647), + ("orderbook_min_samples_per_bucket", 1, 5000), + ("orderbook_min_covered_buckets", 96, 20000), + ("orderbook_min_symbols", 1, 30), ): if key not in result: continue @@ -523,6 +657,8 @@ def _safe_parameters(value: Any) -> dict[str, Any]: result["pooled"] = result["pooled"] is True if "resume_candidate" in result: result["resume_candidate"] = result["resume_candidate"] is True + if "use_orderbook" in result: + result["use_orderbook"] = result["use_orderbook"] is True return result diff --git a/tests/test_orderbook_shadow.py b/tests/test_orderbook_shadow.py new file mode 100644 index 0000000..1948bc8 --- /dev/null +++ b/tests/test_orderbook_shadow.py @@ -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": "", + } diff --git a/tools/calibrate_torch_thresholds.py b/tools/calibrate_torch_thresholds.py index 49b4c36..a689d5c 100644 --- a/tools/calibrate_torch_thresholds.py +++ b/tools/calibrate_torch_thresholds.py @@ -26,6 +26,7 @@ from crypto_spot_bot.bybit import BybitClient from crypto_spot_bot.config import load_settings from crypto_spot_bot.indicators import add_indicators from crypto_spot_bot.models import Candle +from crypto_spot_bot.orderbook_features import ORDERBOOK_FEATURES, load_orderbook_feature_map from crypto_spot_bot.time_series import ( DEFAULT_TORCH_FEATURES, _barrier_outcome, @@ -97,6 +98,14 @@ def main() -> None: context_symbols = sorted(set(symbols + _symbols(args.context_symbols, ()))) horizon = args.horizon if args.horizon > 0 else settings.time_series_forecast_horizon round_trip_cost = _artifact_round_trip_cost(artifact, settings) + orderbook_features: dict[str, dict[int, dict[str, float]]] = {} + if args.orderbook_db: + orderbook_features, _orderbook_manifest = load_orderbook_feature_map( + args.orderbook_db, + interval=settings.base_interval, + symbols=symbols, + min_samples_per_bucket=args.orderbook_min_samples_per_bucket, + ) market_candles: dict[str, list[Candle]] = {} for symbol in context_symbols: @@ -125,6 +134,7 @@ def main() -> None: min_candles=max(30, settings.time_series_min_candles), calibration_window=args.calibration_window, batch_size=args.batch_size, + orderbook_features=orderbook_features, ) records.extend(symbol_records) per_symbol_counts[symbol] = len(symbol_records) @@ -304,6 +314,8 @@ def _parse_args() -> argparse.Namespace: parser.add_argument("--min-oos-folds-with-trades", type=int, default=2, help="Minimum walk-forward folds that must produce trades.") parser.add_argument("--min-oos-profit-factor", type=float, default=1.10, help="Minimum out-of-sample profit factor.") parser.add_argument("--min-benchmark-edge-percent", type=float, default=0.0, help="Required total-net percent advantage over the benchmark.") + parser.add_argument("--orderbook-db", default="", help="SQLite cache used by an artifact with L1 features.") + parser.add_argument("--orderbook-min-samples-per-bucket", type=int, default=20) return parser.parse_args() @@ -352,6 +364,7 @@ def _forecast_records( min_candles: int, calibration_window: int, batch_size: int, + orderbook_features: dict[str, dict[int, dict[str, float]]] | None = None, ) -> list[ForecastRecord]: entry = _torch_recurrent_entry(symbol, artifact) model = _torch_recurrent_model_name(symbol, artifact) @@ -364,6 +377,7 @@ def _forecast_records( symbol=symbol, market_candles=market_candles, trend_candles=trend_candles, + orderbook_features=orderbook_features, ) closes = [float(candle.close) for candle in candles] decision_horizon = _calibration_horizon(entry, horizon, explicit=horizon_is_explicit) @@ -376,6 +390,18 @@ def _forecast_records( start += 1 if calibration_window > 0: start = max(start, end - calibration_window) + lookback = max(1, int(float(entry.get("lookback", 64)))) + requires_orderbook = any(name in ORDERBOOK_FEATURES for name in feature_names) + symbol_orderbook = (orderbook_features or {}).get(symbol.upper(), {}) + valid_indices = { + index + for index in range(start, max(start, end)) + if not requires_orderbook + or all( + candles[position].timestamp in symbol_orderbook + for position in range(index - lookback + 1, index + 1) + ) + } batched_records = _batch_forecast_records( symbol=symbol, candles=candles, @@ -389,6 +415,7 @@ def _forecast_records( start=start, end=end, batch_size=batch_size, + valid_indices=valid_indices, ) if batched_records is not None: return batched_records @@ -398,6 +425,8 @@ def _forecast_records( # belong exclusively to the final quality gate and cannot influence replay. skill = _entry_validation_skill(entry) for index in range(start, max(start, end)): + if index not in valid_indices: + continue prediction = _torch_recurrent_predict( _log_returns(closes[: index + 1]), symbol, @@ -476,6 +505,7 @@ def _batch_forecast_records( start: int, end: int, batch_size: int, + valid_indices: set[int] | None = None, ) -> list[ForecastRecord] | None: if torch is None or RecurrentReturnModel is None: return None @@ -494,7 +524,9 @@ def _batch_forecast_records( indices = [ index for index in range(start, max(start, end)) - if index - lookback + 1 >= 0 and index + decision_horizon < len(closes) + if index - lookback + 1 >= 0 + and index + decision_horizon < len(closes) + and (valid_indices is None or index in valid_indices) ] if not indices: return [] diff --git a/tools/run_torch_retrain.ps1 b/tools/run_torch_retrain.ps1 index 0facd2e..b7c1899 100644 --- a/tools/run_torch_retrain.ps1 +++ b/tools/run_torch_retrain.ps1 @@ -22,6 +22,10 @@ param( [int]$HoldoutWindow = 0, [string]$Interval = "", [string]$EnvFile = "", + [string]$OrderbookDb = "", + [int]$OrderbookMinSamplesPerBucket = 0, + [int]$OrderbookMinCoveredBuckets = 0, + [int]$OrderbookMinSymbols = 0, [switch]$Pooled, [switch]$SkipGuard, [switch]$ResumeCandidate @@ -124,6 +128,10 @@ if ($HoldoutWindow -le 0) { $HoldoutWindow = if ($env:TORCH_RETRAIN_HOLDOUT_WIND if (-not $Interval -and $env:TORCH_RETRAIN_INTERVAL) { $Interval = $env:TORCH_RETRAIN_INTERVAL } if (-not $EnvFile -and $env:TORCH_RETRAIN_ENV) { $EnvFile = $env:TORCH_RETRAIN_ENV } if (-not $EnvFile -and (Test-Path (Join-Path $RepoRoot ".env"))) { $EnvFile = Join-Path $RepoRoot ".env" } +if (-not $OrderbookDb -and $env:TORCH_ORDERBOOK_DB) { $OrderbookDb = $env:TORCH_ORDERBOOK_DB } +if ($OrderbookMinSamplesPerBucket -le 0) { $OrderbookMinSamplesPerBucket = if ($env:TORCH_ORDERBOOK_MIN_SAMPLES_PER_BUCKET) { [int]$env:TORCH_ORDERBOOK_MIN_SAMPLES_PER_BUCKET } else { 20 } } +if ($OrderbookMinCoveredBuckets -le 0) { $OrderbookMinCoveredBuckets = if ($env:TORCH_ORDERBOOK_MIN_COVERED_BUCKETS) { [int]$env:TORCH_ORDERBOOK_MIN_COVERED_BUCKETS } else { 240 } } +if ($OrderbookMinSymbols -le 0) { $OrderbookMinSymbols = if ($env:TORCH_ORDERBOOK_MIN_SYMBOLS) { [int]$env:TORCH_ORDERBOOK_MIN_SYMBOLS } else { 2 } } $ModelFile = if ($env:TIME_SERIES_LSTM_MODEL_PATH) { $env:TIME_SERIES_LSTM_MODEL_PATH } else { Join-Path $RuntimeDir "lstm_forecaster.json" } if (-not [System.IO.Path]::IsPathRooted($ModelFile)) { $ModelFile = Join-Path $RepoRoot $ModelFile } @@ -131,6 +139,10 @@ $CandidateFile = Join-Path $RuntimeDir "lstm_forecaster.candidate.json" $CurrentCalibration = Join-Path $RuntimeDir "torch_guard_current.json" $CandidateCalibration = Join-Path $RuntimeDir "torch_guard_candidate.json" $GuardReport = Join-Path $RuntimeDir "torch_retrain_guard.json" +$ShadowModelFile = Join-Path $RuntimeDir "lstm_forecaster.shadow.json" +$ShadowCalibration = Join-Path $RuntimeDir "torch_shadow_calibration.json" +$ShadowGuard = Join-Path $RuntimeDir "torch_shadow_guard.json" +$ShadowMode = -not [string]::IsNullOrWhiteSpace($OrderbookDb) $mutex = New-Object System.Threading.Mutex($false, "TradeBotTorchRecurrentRetrainer") $hasLock = $false @@ -177,6 +189,14 @@ try { if ($Features) { $trainerArgs += @("--features", $Features) } if ($ContextSymbols) { $trainerArgs += @("--context-symbols", $ContextSymbols) } if ($Seed -gt 0) { $trainerArgs += @("--seed", $Seed.ToString()) } + if ($OrderbookDb) { + $trainerArgs += @( + "--orderbook-db", $OrderbookDb, + "--orderbook-min-samples-per-bucket", $OrderbookMinSamplesPerBucket.ToString(), + "--orderbook-min-covered-buckets", $OrderbookMinCoveredBuckets.ToString(), + "--orderbook-min-symbols", $OrderbookMinSymbols.ToString() + ) + } Push-Location $RepoRoot $pushedLocation = $true @@ -216,6 +236,12 @@ try { ) if ($Symbols) { $calibrationBaseArgs += @("--symbols", $Symbols) } if ($EnvFile) { $calibrationBaseArgs += @("--env", $EnvFile) } + if ($OrderbookDb) { + $calibrationBaseArgs += @( + "--orderbook-db", $OrderbookDb, + "--orderbook-min-samples-per-bucket", $OrderbookMinSamplesPerBucket.ToString() + ) + } if (Test-Path $ModelFile) { Write-RetrainLog "Calibrating current artifact for guard." @@ -243,13 +269,14 @@ try { } Write-RetrainLog "Running retrain guard." + $GuardTarget = if ($ShadowMode) { $ShadowModelFile } else { $ModelFile } $guardArgs = @( "-u", "tools\accept_torch_candidate.py", "--current-report", $CurrentCalibration, "--candidate-report", $CandidateCalibration, "--candidate-artifact", $CandidateFile, - "--target-artifact", $ModelFile, + "--target-artifact", $GuardTarget, "--report", $GuardReport ) $guardExitCode = Invoke-LoggedNativeCommand -FilePath $python -ArgumentList $guardArgs -LogPath $LogFile @@ -261,10 +288,22 @@ try { throw "Retrain guard failed with exit code $guardExitCode." } if (Test-Path $CandidateCalibration) { - Copy-Item -Force -LiteralPath $CandidateCalibration -Destination (Join-Path $RuntimeDir "torch_threshold_calibration.json") - Write-RetrainLog "Updated active threshold calibration: $(Join-Path $RuntimeDir "torch_threshold_calibration.json")" + if ($ShadowMode) { + Copy-Item -Force -LiteralPath $CandidateCalibration -Destination $ShadowCalibration + Copy-Item -Force -LiteralPath $GuardReport -Destination $ShadowGuard + Write-RetrainLog "Candidate passed offline gate and was staged for shadow only: $ShadowModelFile" + } + else { + Copy-Item -Force -LiteralPath $CandidateCalibration -Destination (Join-Path $RuntimeDir "torch_threshold_calibration.json") + Write-RetrainLog "Updated active threshold calibration: $(Join-Path $RuntimeDir "torch_threshold_calibration.json")" + } + } + if ($ShadowMode) { + Write-RetrainLog "Candidate accepted by offline guard. Active artifact was not changed: $ModelFile" + } + else { + Write-RetrainLog "Candidate accepted by guard. Active artifact: $ModelFile" } - Write-RetrainLog "Candidate accepted by guard. Active artifact: $ModelFile" } catch { Write-RetrainLog "ERROR: $($_.Exception.Message)" diff --git a/tools/sync_orderbook_observations.py b/tools/sync_orderbook_observations.py new file mode 100644 index 0000000..e378a61 --- /dev/null +++ b/tools/sync_orderbook_observations.py @@ -0,0 +1,210 @@ +from __future__ import annotations + +import argparse +import base64 +import json +import os +import sqlite3 +from pathlib import Path +from typing import Any +from urllib.error import HTTPError, URLError +from urllib.parse import urlencode +from urllib.request import Request, urlopen + + +def sync_orderbook_observations( + *, + api_base_url: str, + token: str, + database_path: str | Path, + timeout: int = 60, + page_limit: int = 5000, +) -> dict[str, Any]: + path = Path(database_path) + path.parent.mkdir(parents=True, exist_ok=True) + _init_schema(path) + manifest = _get_json( + api_base_url, + "/api/training/market-observations/manifest", + token=token, + timeout=timeout, + ) + rows = manifest.get("items") if isinstance(manifest.get("items"), list) else [] + downloaded = 0 + symbol_results: list[dict[str, Any]] = [] + for row in rows: + if not isinstance(row, dict): + continue + symbol = str(row.get("symbol") or "").strip().upper() + remote_max_id = int(row.get("max_id", 0) or 0) + if not symbol or remote_max_id <= 0: + continue + after_id = _local_max_id(path, symbol) + symbol_downloaded = 0 + while after_id < remote_max_id: + query = urlencode( + { + "symbol": symbol, + "after_id": after_id, + "limit": max(1, min(5000, int(page_limit))), + } + ) + payload = _get_json( + api_base_url, + f"/api/training/market-observations?{query}", + token=token, + timeout=timeout, + ) + items = payload.get("items") if isinstance(payload.get("items"), list) else [] + if not items: + break + inserted = _insert_rows(path, items) + symbol_downloaded += inserted + downloaded += inserted + next_after_id = int(payload.get("next_after_id", after_id) or after_id) + if next_after_id <= after_id: + break + after_id = next_after_id + symbol_results.append( + { + "symbol": symbol, + "downloaded": symbol_downloaded, + "local_max_id": _local_max_id(path, symbol), + "remote_max_id": remote_max_id, + } + ) + return { + "database_path": str(path.resolve()), + "downloaded": downloaded, + "symbols": symbol_results, + "local_samples": _local_count(path), + } + + +def _init_schema(path: Path) -> None: + with sqlite3.connect(path) as connection: + connection.executescript( + """ + PRAGMA journal_mode=WAL; + CREATE TABLE IF NOT EXISTS market_observations ( + id INTEGER PRIMARY KEY, + symbol TEXT NOT NULL, + bid_price REAL NOT NULL, + bid_size REAL NOT NULL, + ask_price REAL NOT NULL, + ask_size REAL NOT NULL, + mid_price REAL NOT NULL, + microprice REAL NOT NULL, + spread_bps REAL NOT NULL, + imbalance REAL NOT NULL, + last_price REAL NOT NULL, + source_timestamp_ms INTEGER NOT NULL DEFAULT 0, + created_at TEXT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_local_market_observations_symbol_id + ON market_observations(symbol, id); + """ + ) + + +def _insert_rows(path: Path, rows: list[Any]) -> int: + values = [] + for row in rows: + if not isinstance(row, dict): + continue + values.append( + ( + int(row.get("id", 0) or 0), + str(row.get("symbol") or "").upper(), + float(row.get("bid_price", 0.0) or 0.0), + float(row.get("bid_size", 0.0) or 0.0), + float(row.get("ask_price", 0.0) or 0.0), + float(row.get("ask_size", 0.0) or 0.0), + float(row.get("mid_price", 0.0) or 0.0), + float(row.get("microprice", 0.0) or 0.0), + float(row.get("spread_bps", 0.0) or 0.0), + float(row.get("imbalance", 0.0) or 0.0), + float(row.get("last_price", 0.0) or 0.0), + int(row.get("source_timestamp_ms", 0) or 0), + str(row.get("created_at") or ""), + ) + ) + if not values: + return 0 + with sqlite3.connect(path) as connection: + before = connection.total_changes + connection.executemany( + """ + INSERT OR IGNORE INTO market_observations ( + id, symbol, bid_price, bid_size, ask_price, ask_size, + mid_price, microprice, spread_bps, imbalance, last_price, + source_timestamp_ms, created_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + values, + ) + return connection.total_changes - before + + +def _local_max_id(path: Path, symbol: str) -> int: + with sqlite3.connect(path) as connection: + row = connection.execute( + "SELECT MAX(id) FROM market_observations WHERE symbol = ?", + (symbol,), + ).fetchone() + return int(row[0] or 0) if row else 0 + + +def _local_count(path: Path) -> int: + with sqlite3.connect(path) as connection: + row = connection.execute("SELECT COUNT(*) FROM market_observations").fetchone() + return int(row[0] or 0) if row else 0 + + +def _get_json(api_base_url: str, path: str, *, token: str, timeout: int) -> dict[str, Any]: + headers = {"Accept": "application/json"} + headers.update(_auth_headers(token)) + request = Request(api_base_url.rstrip("/") + path, headers=headers, method="GET") + try: + with urlopen(request, timeout=timeout) as response: + text = response.read().decode("utf-8") + except HTTPError as exc: + detail = exc.read().decode("utf-8", errors="replace") + raise RuntimeError(f"HTTP {exc.code} {path}: {detail[:300]}") from exc + except URLError as exc: + raise RuntimeError(f"network error {path}: {exc.reason}") from exc + data = json.loads(text) if text.strip() else {} + return data if isinstance(data, dict) else {} + + +def _auth_headers(token: str) -> dict[str, str]: + value = token.strip() + if not value: + return {} + headers = {"X-TradeBot-Token": value} + if value.lower().startswith(("basic ", "bearer ")): + headers["Authorization"] = value + elif ":" in value: + encoded = base64.b64encode(value.encode("utf-8")).decode("ascii") + headers["Authorization"] = f"Basic {encoded}" + else: + headers["Authorization"] = f"Bearer {value}" + return headers + + +def main() -> None: + parser = argparse.ArgumentParser(description="Synchronize TradeBot L1 observations to a local SQLite cache.") + parser.add_argument("--api-base-url", default=os.environ.get("TRADEBOT_API_BASE_URL", "https://tb.kusoft.xyz")) + parser.add_argument("--api-auth", default=os.environ.get("TRADEBOT_API_AUTH", "")) + parser.add_argument("--database", default="runtime/orderbook_observations.sqlite3") + args = parser.parse_args() + result = sync_orderbook_observations( + api_base_url=args.api_base_url, + token=args.api_auth, + database_path=args.database, + ) + print(json.dumps(result, ensure_ascii=False)) + + +if __name__ == "__main__": + main() diff --git a/tools/train_torch_recurrent_forecaster.py b/tools/train_torch_recurrent_forecaster.py index 2000bfa..fba59a7 100644 --- a/tools/train_torch_recurrent_forecaster.py +++ b/tools/train_torch_recurrent_forecaster.py @@ -28,6 +28,7 @@ from crypto_spot_bot.bybit import BybitClient from crypto_spot_bot.config import load_settings from crypto_spot_bot.indicators import add_indicators from crypto_spot_bot.models import Candle +from crypto_spot_bot.orderbook_features import ORDERBOOK_FEATURES, load_orderbook_feature_map from crypto_spot_bot.time_series import ( DEFAULT_TORCH_FEATURES, _barrier_outcome, @@ -41,6 +42,8 @@ EVENT_OUTPUT_NAME = "logit_tp_first" OUTPUT_LAYOUT = (*RETURN_OUTPUT_LAYOUT, EVENT_OUTPUT_NAME) TARGET_TRANSFORM = "barrier_net_return" QUANTILES = {"q10": 0.10, "q50": 0.50, "q90": 0.90} +_ORDERBOOK_FEATURES_BY_SYMBOL: dict[str, dict[int, dict[str, float]]] = {} +_ORDERBOOK_MANIFEST: dict[str, dict[str, Any]] = {} @dataclass(slots=True) @@ -161,6 +164,7 @@ class RecurrentReturnModel(nn.Module): def main() -> None: + global _ORDERBOOK_FEATURES_BY_SYMBOL, _ORDERBOOK_MANIFEST args = _parse_args() if args.threads > 0: torch.set_num_threads(args.threads) @@ -175,6 +179,33 @@ def main() -> None: decision_horizon = args.horizon if args.horizon > 0 else max(1, settings.time_series_forecast_horizon) target_horizons = _horizons(args.horizons, decision_horizon) feature_names = _feature_names_arg(args.features) + if args.orderbook_db: + _ORDERBOOK_FEATURES_BY_SYMBOL, _ORDERBOOK_MANIFEST = load_orderbook_feature_map( + args.orderbook_db, + interval=interval, + symbols=symbols, + min_samples_per_bucket=args.orderbook_min_samples_per_bucket, + ) + eligible_symbols = [ + symbol + for symbol in symbols + if int(_ORDERBOOK_MANIFEST.get(symbol, {}).get("covered_buckets", 0) or 0) + >= args.orderbook_min_covered_buckets + ] + if len(eligible_symbols) < max(1, args.orderbook_min_symbols): + coverage = ", ".join( + f"{symbol}={int(_ORDERBOOK_MANIFEST.get(symbol, {}).get('covered_buckets', 0) or 0)}" + for symbol in symbols + ) + raise SystemExit( + "Orderbook coverage is below the training minimum: " + f"need {args.orderbook_min_covered_buckets} buckets for " + f"{args.orderbook_min_symbols} symbols; got {coverage or 'no data'}" + ) + symbols = eligible_symbols + for feature_name in ORDERBOOK_FEATURES: + if feature_name not in feature_names: + feature_names.append(feature_name) if args.pooled: feature_names.extend(f"symbol_is_{symbol}" for symbol in symbols) ensemble_seeds = _ints(args.ensemble_seeds) or [args.seed] @@ -214,6 +245,15 @@ def main() -> None: "selection_folds": args.selection_folds, "symbols": {}, } + if args.orderbook_db: + artifact["orderbook_features"] = { + "source": "forward_collected_bybit_l1", + "interval": interval, + "min_samples_per_bucket": args.orderbook_min_samples_per_bucket, + "min_covered_buckets": args.orderbook_min_covered_buckets, + "features": list(ORDERBOOK_FEATURES), + "coverage": {symbol: _ORDERBOOK_MANIFEST.get(symbol, {}) for symbol in symbols}, + } if args.pooled: artifact["version"] = 7 @@ -584,6 +624,10 @@ def _parse_args() -> argparse.Namespace: parser.add_argument("--threads", type=int, default=0, help="Torch CPU threads; 0 keeps torch default.") parser.add_argument("--device", default="auto", help="auto, cpu, cuda, or mps.") parser.add_argument("--output", default="", help="Output JSON path. Defaults to TIME_SERIES_LSTM_MODEL_PATH.") + parser.add_argument("--orderbook-db", default="", help="SQLite cache containing forward-collected L1 observations.") + parser.add_argument("--orderbook-min-samples-per-bucket", type=int, default=20) + parser.add_argument("--orderbook-min-covered-buckets", type=int, default=240) + parser.add_argument("--orderbook-min-symbols", type=int, default=2) return parser.parse_args() @@ -802,6 +846,7 @@ def _prepare_data( symbol=symbol, market_candles=market_candles, trend_candles=trend_candles, + orderbook_features=_ORDERBOOK_FEATURES_BY_SYMBOL, ) max_horizon = max(target_horizons) samples: list[TrainingSample] = [] @@ -812,6 +857,11 @@ def _prepare_data( window = feature_rows[end_index - lookback + 1 : end_index + 1] if len(window) != lookback: continue + if any(name in ORDERBOOK_FEATURES for name in feature_names): + symbol_orderbook = _ORDERBOOK_FEATURES_BY_SYMBOL.get(symbol.upper(), {}) + window_candles = candles[end_index - lookback + 1 : end_index + 1] + if any(row.timestamp not in symbol_orderbook for row in window_candles): + continue raw_targets: list[float] = [] event_targets: list[float] = [] volatility_scales: list[float] = [] @@ -856,6 +906,8 @@ def _prepare_data( validation_window = min(max(16, validation_window), max(16, validation_end // 3)) validation_start = validation_end - validation_window train_end = validation_start - max_horizon + if validation_start < 0 or train_end <= 0: + return None train_samples = samples[:train_end] validation_samples = samples[validation_start:validation_end] holdout_samples = samples[holdout_start:] diff --git a/tools/windows_training_agent.py b/tools/windows_training_agent.py index 4faf351..6a2c766 100644 --- a/tools/windows_training_agent.py +++ b/tools/windows_training_agent.py @@ -20,12 +20,24 @@ from urllib.error import URLError from urllib.request import Request from urllib.request import urlopen +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from crypto_spot_bot.orderbook_features import load_orderbook_feature_map +from tools.sync_orderbook_observations import sync_orderbook_observations + ARTIFACT_NAMES = ( "lstm_forecaster.json", "torch_retrain_guard.json", "torch_threshold_calibration.json", ) +SHADOW_ARTIFACT_NAMES = ( + "lstm_forecaster.shadow.json", + "torch_shadow_guard.json", + "torch_shadow_calibration.json", +) def main() -> None: @@ -65,7 +77,35 @@ def poll_once(args: argparse.Namespace, repo_root: Path, runtime_dir: Path, log_ message = "" summary: dict[str, Any] = {} try: - run_retrain(args, job_id, lease_token, job, repo_root, log_path) + parameters = job.get("parameters") if isinstance(job.get("parameters"), dict) else {} + use_orderbook = parameters.get("use_orderbook", True) is not False + orderbook_status: dict[str, Any] = {} + if use_orderbook: + report_progress( + args, + job_id, + lease_token, + "running", + "orderbook_sync", + 4, + "Синхронизирую forward-наблюдения стакана", + ) + orderbook_status = prepare_orderbook_data(args, repo_root, parameters, log_path) + if orderbook_status["state"] != "ready": + summary = orderbook_status + message = "forward orderbook coverage is still accumulating" + success = True + log(log_path, f"Job {job_id} remains in collecting state: {orderbook_status}") + return + run_retrain( + args, + job_id, + lease_token, + job, + repo_root, + log_path, + orderbook_db=(repo_root / "runtime" / "orderbook_observations.sqlite3") if use_orderbook else None, + ) summary = read_json(runtime_dir / "torch_retrain_guard.json") accepted = summary.get("accepted") is True if accepted: @@ -78,12 +118,20 @@ def poll_once(args: argparse.Namespace, repo_root: Path, runtime_dir: Path, log_ 72, "Обучение завершено, загружаю артефакты", ) - for name in ARTIFACT_NAMES: + artifact_names = SHADOW_ARTIFACT_NAMES if use_orderbook else ARTIFACT_NAMES + if use_orderbook: + summary["deployment"] = "shadow" + summary["orderbook"] = orderbook_status + for name in artifact_names: path = runtime_dir / name if path.is_file(): upload_artifact(args, job_id, lease_token, path, log_path) - message = "training completed; candidate accepted" - log(log_path, f"Completed retrain job {job_id}; candidate accepted") + message = ( + "training completed; candidate staged in shadow" + if use_orderbook + else "training completed; candidate accepted" + ) + log(log_path, f"Completed retrain job {job_id}; {message}") else: reason = str(summary.get("reason") or "validation failed") message = f"training completed; candidate rejected by quality gate: {reason}" @@ -109,6 +157,7 @@ def run_retrain( job: dict[str, Any], repo_root: Path, log_path: Path, + orderbook_db: Path | None = None, ) -> None: script = repo_root / "tools" / "run_torch_retrain.ps1" if not script.is_file(): @@ -153,6 +202,14 @@ def run_retrain( cmd.append("-Pooled") if parameters.get("resume_candidate") is True: cmd.append("-ResumeCandidate") + if orderbook_db is not None: + cmd.extend(["-OrderbookDb", str(orderbook_db)]) + for key, ps_arg, default in ( + ("orderbook_min_samples_per_bucket", "-OrderbookMinSamplesPerBucket", 20), + ("orderbook_min_covered_buckets", "-OrderbookMinCoveredBuckets", 240), + ("orderbook_min_symbols", "-OrderbookMinSymbols", 2), + ): + cmd.extend([ps_arg, str(int(parameters.get(key, default) or default))]) log(log_path, "Running retrain: " + " ".join(quote_for_log(part) for part in cmd)) report_progress( args, @@ -236,6 +293,65 @@ def run_retrain( ) +def prepare_orderbook_data( + args: argparse.Namespace, + repo_root: Path, + parameters: dict[str, Any], + log_path: Path, +) -> dict[str, Any]: + database_path = repo_root / "runtime" / "orderbook_observations.sqlite3" + token = args.api_auth or os.environ.get("TRADEBOT_API_AUTH", "") + sync_result = sync_orderbook_observations( + api_base_url=args.api_base_url, + token=token, + database_path=database_path, + ) + interval = str(parameters.get("interval") or os.environ.get("TORCH_RETRAIN_INTERVAL") or "60") + minimum_samples = int(parameters.get("orderbook_min_samples_per_bucket", 20) or 20) + minimum_buckets = int(parameters.get("orderbook_min_covered_buckets", 240) or 240) + minimum_symbols = int(parameters.get("orderbook_min_symbols", 2) or 2) + requested_symbols = { + item.strip().upper() + for item in str(parameters.get("symbols") or "").split(",") + if item.strip() + } + _features, manifest = load_orderbook_feature_map( + database_path, + interval=interval, + symbols=sorted(requested_symbols) if requested_symbols else None, + min_samples_per_bucket=minimum_samples, + ) + eligible = sorted( + symbol + for symbol, row in manifest.items() + if int(row.get("covered_buckets", 0) or 0) >= minimum_buckets + ) + state = "ready" if len(eligible) >= minimum_symbols else "collecting_orderbook" + coverage = { + symbol: int(row.get("covered_buckets", 0) or 0) + for symbol, row in sorted(manifest.items()) + } + result = { + "accepted": False, + "state": state, + "reason": ( + "orderbook coverage ready for training" + if state == "ready" + else "forward orderbook coverage is below the configured minimum" + ), + "eligible_symbols": eligible, + "eligible_symbol_count": len(eligible), + "minimum_symbols": minimum_symbols, + "minimum_covered_buckets": minimum_buckets, + "minimum_samples_per_bucket": minimum_samples, + "covered_buckets_by_symbol": coverage, + "local_samples": int(sync_result.get("local_samples", 0) or 0), + "downloaded_samples": int(sync_result.get("downloaded", 0) or 0), + } + log(log_path, "Orderbook preparation: " + json.dumps(result, ensure_ascii=False, sort_keys=True)) + return result + + def friendly_training_message(message: str) -> str: cleaned = message.strip() if not cleaned: