diff --git a/.env.example b/.env.example index 65ebbb2..6a83b51 100644 --- a/.env.example +++ b/.env.example @@ -72,6 +72,9 @@ TIME_SERIES_PROBE_MIN_EDGE_PERCENT=0.02 TIME_SERIES_PROBE_MIN_PROBABILITY_UP=0.55 TIME_SERIES_PROBE_SIZE_MULTIPLIER=0.40 TIME_SERIES_REBOUND_FALLBACK_ENABLED=false +# Use the independently guarded trend/MACD strategy while no accepted fresh +# Torch model is available. The rejected model is never used for entries. +TIME_SERIES_TREND_FALLBACK_ENABLED=true TIME_SERIES_REQUIRE_QUALITY_GATE=true # Emergency paper-only override. Keep false unless a failed guard is accepted manually. TIME_SERIES_MANUAL_QUALITY_OVERRIDE=false @@ -107,7 +110,7 @@ STORAGE_RETENTION_DAYS=30 STORAGE_PRUNE_INTERVAL_SECONDS=3600 # Windows trainer keeps this final tail untouched by training and early stopping. -TORCH_RETRAIN_HOLDOUT_WINDOW=240 +TORCH_RETRAIN_HOLDOUT_WINDOW=1000 DATABASE_PATH=runtime/tradebot.sqlite3 LOG_PATH=runtime/tradebot.log diff --git a/README.md b/README.md index f9d12e3..8d9e5f8 100644 --- a/README.md +++ b/README.md @@ -10,6 +10,7 @@ Spot-бот для демо-торговли криптовалютой на р - Spot-only логика: покупка базовой монеты за USDT и продажа обратно, без short и без плеча. - Live spot-ордеры явно отправляются без плеча: `category=spot`, `isLeverage=0`. - Основная стратегия `torch_forecast`: входы и forecast-выходы идут только от свежей экспортированной PyTorch LSTM/GRU модели с успешным quality gate; MACD/RSI/дневная EMA не являются условиями входа в этом режиме. Rebound fallback без модели выключен по умолчанию. Спред, ликвидность, stop-loss, ATR trailing stop, запрет DCA и лимиты экспозиции остаются защитой исполнения и риска. +- При `TIME_SERIES_TREND_FALLBACK_ENABLED=true` отсутствие принятой свежей Torch-модели включает самостоятельную `trend_macd`-стратегию. Отклонённый artifact не используется, fallback явно отражается в readiness и диагностике сигналов, а после появления принятой модели выключается автоматически. - Основная стратегия `trend_macd`: вход на `1h`, дневной фильтр тренда на `1d`, long только если цена выше дневной EMA200 и дневная EMA50 выше EMA200. - Вход `trend_macd`: MACD на `1h` пересекает signal вверх, цена выше EMA50, RSI в диапазоне `45..65`, спред и ликвидность проходят runtime-фильтры. - Выход `trend_macd`: MACD пересекает signal вниз, `1h` свеча закрылась ниже EMA50, сработал стоп `4%` или ATR trailing stop `2.2 ATR`. @@ -81,7 +82,7 @@ Dashboard: Новый artifact версии 4 обучается как probabilistic multi-horizon модель: вход включает доходности, форму свечи, объем, ATR%, realized volatility, RSI/MACD/EMA slopes, 4h/24h rolling trend, дневные EMA-признаки, BTC/ETH cross-asset признаки и числовые признаки текущего шаблона пары. Цель обучается как `future log return - комиссии - проскальзывание`, нормализованная на текущую волатильность. Модель сразу прогнозирует горизонты `1/3/6/12`, quantile-оценки `q10/q50/q90` и `P(up)`. -Последний tail (`--holdout-window`, по умолчанию 240 samples на символ) полностью исключается из training и early stopping. Между train/validation/holdout оставляется purge по максимальному forecast horizon. Threshold walk-forward и guard работают только на этом untouched holdout; calibration и guard криптографически привязаны к SHA-256 конкретного model artifact. +Последний tail (`--holdout-window`, по умолчанию 1000 samples на символ) полностью исключается из training и early stopping. Между train/validation/holdout оставляется purge по максимальному forecast horizon. Threshold walk-forward и guard работают только на этом untouched holdout; calibration и guard криптографически привязаны к SHA-256 конкретного model artifact. В каждом walk-forward fold торговать могут только пары, которые получили жизнеспособный порог на предшествующей train-части; общий порог больше не возвращает в портфель нестабильные пары. Файл из `TIME_SERIES_LSTM_MODEL_PATH` читается ботом автоматически, если `TIME_SERIES_FORECAST_ENABLED=true`. В стратегии `torch_forecast` экспортированная PyTorch LSTM/GRU модель является единственным направляющим сигналом для входа и forecast-выхода. Экспортированные модели появляются в dashboard как `PyTorch LSTM` или `PyTorch GRU`; старый легкий reservoir LSTM-кандидат и все встроенные не-torch прогнозы удалены. @@ -100,9 +101,9 @@ powershell -ExecutionPolicy Bypass -File tools\install_windows_training_agent.ps Установщик сохраняет worker-токен через Windows DPAPI, удаляет его старую plaintext-копию из пользовательского окружения и включает постоянный запуск агента. С правами администратора используется Scheduled Task с watchdog; без повышения прав — штатный ярлык в пользовательской папке Startup. Старые локальные retrain-задачи удаляются, чтобы обучение запускалось через очередь, а не двумя независимыми механизмами. -По умолчанию Windows-agent обучает pooled multi-asset PyTorch `LSTM/GRU` на `6000` часовых свечах: общие recurrent-веса получают one-hot embedding символа, прогноз усредняется по seed `7/19/43`, модели сравниваются на validation-folds, а пороги калибруются отдельно для каждой пары. Search space использует lookback `32/64/128`, hidden `64/96`, dropout `0.20`, AdamW learning rate `0.0007` и weight decay `0.0005`; untouched holdout и quality gate не ослабляются. Параметры можно переопределить через env: `TORCH_RETRAIN_SYMBOLS`, `TORCH_RETRAIN_LIMIT`, `TORCH_RETRAIN_LOOKBACKS`, `TORCH_RETRAIN_ARCHITECTURES`, `TORCH_RETRAIN_HIDDEN_SIZES`, `TORCH_RETRAIN_LAYERS`, `TORCH_RETRAIN_DROPOUTS`, `TORCH_RETRAIN_HORIZON`, `TORCH_RETRAIN_HORIZONS`, `TORCH_RETRAIN_CONTEXT_SYMBOLS`, `TORCH_RETRAIN_FEATURES`, `TORCH_RETRAIN_SEED`, `TORCH_RETRAIN_ENSEMBLE_SEEDS`, `TORCH_RETRAIN_SELECTION_FOLDS`, `TORCH_RETRAIN_LEARNING_RATE`, `TORCH_RETRAIN_WEIGHT_DECAY`, `TORCH_RETRAIN_EPOCHS`, `TORCH_RETRAIN_PATIENCE`, `TORCH_RETRAIN_INTERVAL`, `TORCH_RETRAIN_ENV`. +По умолчанию Windows-agent обучает отдельную PyTorch `LSTM/GRU` для каждой пары на `6000` часовых свечах. Это не заставляет разнородные активы делить одну архитектуру и один набор recurrent-весов. Прогноз усредняется по seed `7/19`, модели сравниваются на validation-folds, а пороги калибруются отдельно для каждой пары. Ensemble guard выполняется пакетно на GPU, а экспорт не дублирует первый набор весов. Search space использует lookback `32/64/128`, hidden `64/96`, dropout `0.20`, AdamW learning rate `0.0007` и weight decay `0.0005`; untouched holdout и quality gate не ослабляются. Для диагностического pooled-запуска используется ключ `-Pooled`. Параметры можно переопределить через env: `TORCH_RETRAIN_SYMBOLS`, `TORCH_RETRAIN_LIMIT`, `TORCH_RETRAIN_LOOKBACKS`, `TORCH_RETRAIN_ARCHITECTURES`, `TORCH_RETRAIN_HIDDEN_SIZES`, `TORCH_RETRAIN_LAYERS`, `TORCH_RETRAIN_DROPOUTS`, `TORCH_RETRAIN_HORIZON`, `TORCH_RETRAIN_HORIZONS`, `TORCH_RETRAIN_CONTEXT_SYMBOLS`, `TORCH_RETRAIN_FEATURES`, `TORCH_RETRAIN_SEED`, `TORCH_RETRAIN_ENSEMBLE_SEEDS`, `TORCH_RETRAIN_SELECTION_FOLDS`, `TORCH_RETRAIN_LEARNING_RATE`, `TORCH_RETRAIN_WEIGHT_DECAY`, `TORCH_RETRAIN_EPOCHS`, `TORCH_RETRAIN_PATIENCE`, `TORCH_RETRAIN_INTERVAL`, `TORCH_RETRAIN_ENV`. -Loss и выбор гиперпараметров учитывают after-cost trading utility и ранговую связь прогноза с будущей доходностью, а не только MAE. В каждом walk-forward fold вероятность `P(up)` калибруется Platt-моделью исключительно на train-части; затем на этой же train-части выбираются глобальные и per-symbol пороги, которые применяются к test-части. Калибратор не имеет fallback на единичные сделки: если минимальная статистика не набрана, кандидат получает `calibration_insufficient` и не может пройти gate. +Loss и выбор гиперпараметров учитывают after-cost trading utility и ранговую связь прогноза с будущей доходностью, а не только MAE. В каждом walk-forward fold вероятность `P(up)` калибруется Platt-моделью исключительно на train-части; затем на этой же train-части выбираются глобальные и per-symbol пороги, которые применяются к test-части. Для выбора порога требуется минимум 24 непересекающиеся сделки, а финальный quality gate по-прежнему требует не менее 30 OOS-сделок. Калибратор не имеет fallback на единичные сделки: если минимальная статистика не набрана, кандидат получает `calibration_insufficient` и не может пройти gate. Основной decision horizon — `12h`, дополнительные горизонты — `3/6/12/24`. Это согласует прогноз с round-trip cost: при текущих fee/slippage полный вход-выход стоит около `0.26%`, поэтому прежний `3h` target чаще описывал шум, который не покрывал издержки. Threshold search оценивается тем же execution replay со stop-loss, take-profit, ATR trailing и forecast-exit, который используется в walk-forward. `holdout_skill` остаётся только в финальном отчёте и никогда не участвует в фильтрации входов или подборе порогов. diff --git a/crypto_spot_bot/bot.py b/crypto_spot_bot/bot.py index 7c2ab8d..ceb8952 100644 --- a/crypto_spot_bot/bot.py +++ b/crypto_spot_bot/bot.py @@ -12,7 +12,7 @@ from crypto_spot_bot.learning import TradeLearner from crypto_spot_bot.market_data import MarketData 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 +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 @@ -78,6 +78,9 @@ class CryptoSpotBot: self.started_at = utc_now() self.message = "бот работает" self._safe_event("Бот запущен") + # Maintenance must never delay the first market decision after startup. + # The bounded telemetry prune starts after the configured interval. + self._last_prune_at = utc_now() if self.settings.websocket_enabled: self._ws_task = asyncio.create_task(self.market.websocket_loop()) self._loop_task = asyncio.create_task(self._run_loop()) @@ -458,20 +461,19 @@ class CryptoSpotBot: invalid_models = [] for symbol in self.market.symbols: forecast = self.market.forecasts.get(symbol, {}) - if not forecast.get("usable"): - invalid_models.append(symbol) - continue - if ( - self.settings.time_series_require_quality_gate - and not self.settings.time_series_manual_quality_override - and forecast.get("quality_gate_passed") is not True - ): - invalid_models.append(symbol) - continue - if self.settings.time_series_require_fresh_model and forecast.get("model_fresh") is not True: + if torch_model_readiness_reasons(self.settings, forecast): invalid_models.append(symbol) if invalid_models: - reasons.append("forecast_model_not_ready") + if self.settings.time_series_trend_fallback_enabled: + forecast_fallback_active = True + else: + forecast_fallback_active = False + reasons.append("forecast_model_not_ready") + else: + forecast_fallback_active = False + else: + invalid_models = [] + forecast_fallback_active = False reconciliation: dict = {} if isinstance(self.broker, LiveBroker): reconciliation = dict(self.broker.reconciliation_state) @@ -485,6 +487,9 @@ class CryptoSpotBot: "stale_symbols": stale_symbols, "consecutive_loop_errors": self._consecutive_loop_errors, "reconciliation": reconciliation, + "forecast_model_ready": not invalid_models, + "forecast_fallback_active": forecast_fallback_active, + "forecast_invalid_symbols": invalid_models, } def account_snapshot(self) -> dict[str, float]: diff --git a/crypto_spot_bot/bybit.py b/crypto_spot_bot/bybit.py index 02e0aaa..5462176 100644 --- a/crypto_spot_bot/bybit.py +++ b/crypto_spot_bot/bybit.py @@ -42,7 +42,11 @@ class Instrument: class BybitClient: def __init__(self, settings: Settings): self.settings = settings - self.session = requests.Session() + self.session = self._build_session() + + @staticmethod + def _build_session() -> requests.Session: + session = requests.Session() retry = Retry( total=3, connect=3, @@ -53,14 +57,32 @@ class BybitClient: allowed_methods=frozenset({"GET"}), respect_retry_after_header=True, ) - self.session.mount("https://", HTTPAdapter(max_retries=retry)) + session.mount("https://", HTTPAdapter(max_retries=retry)) + return session + + def _reset_session(self) -> None: + self.session.close() + self.session = self._build_session() def public_get(self, path: str, params: dict[str, Any]) -> dict[str, Any]: - response = self.session.get( - f"{self.settings.rest_base_url}{path}", - params=params, - timeout=12, - ) + response = None + for attempt in range(3): + try: + response = self.session.get( + f"{self.settings.rest_base_url}{path}", + params=params, + timeout=12, + ) + break + except (requests.exceptions.ConnectionError, requests.exceptions.Timeout): + if attempt >= 2: + raise + # A failed TLS session can remain poisoned in urllib3's pool. + # Recreate the pool before retrying instead of reusing it. + self._reset_session() + time.sleep(0.5 * (2**attempt)) + if response is None: # pragma: no cover - loop either returns or raises. + raise BybitError("Bybit public request produced no response") response.raise_for_status() return self._unwrap(response.json()) diff --git a/crypto_spot_bot/config.py b/crypto_spot_bot/config.py index a4ef365..5737e92 100644 --- a/crypto_spot_bot/config.py +++ b/crypto_spot_bot/config.py @@ -137,6 +137,7 @@ class Settings: time_series_probe_min_probability_up: float time_series_probe_size_multiplier: float time_series_rebound_fallback_enabled: bool + time_series_trend_fallback_enabled: bool stop_loss_percent: float stop_loss_exit_enabled: bool take_profit_percent: float @@ -304,6 +305,7 @@ def load_settings(env_file: str | Path | None = None) -> Settings: time_series_probe_min_probability_up=_float_env("TIME_SERIES_PROBE_MIN_PROBABILITY_UP", 0.55), time_series_probe_size_multiplier=_float_env("TIME_SERIES_PROBE_SIZE_MULTIPLIER", 0.40), time_series_rebound_fallback_enabled=_bool_env("TIME_SERIES_REBOUND_FALLBACK_ENABLED", False), + time_series_trend_fallback_enabled=_bool_env("TIME_SERIES_TREND_FALLBACK_ENABLED", False), stop_loss_percent=_float_env("STOP_LOSS_PERCENT", 0.04), stop_loss_exit_enabled=_bool_env("STOP_LOSS_EXIT_ENABLED", True), take_profit_percent=_float_env("TAKE_PROFIT_PERCENT", 0.035), diff --git a/crypto_spot_bot/dashboard.py b/crypto_spot_bot/dashboard.py index eb5d172..7d9e4b7 100644 --- a/crypto_spot_bot/dashboard.py +++ b/crypto_spot_bot/dashboard.py @@ -398,6 +398,7 @@ def _safe_config(settings: Settings) -> dict[str, Any]: "time_series_probe_min_probability_up": settings.time_series_probe_min_probability_up, "time_series_probe_size_multiplier": settings.time_series_probe_size_multiplier, "time_series_rebound_fallback_enabled": settings.time_series_rebound_fallback_enabled, + "time_series_trend_fallback_enabled": settings.time_series_trend_fallback_enabled, "time_series_require_quality_gate": settings.time_series_require_quality_gate, "time_series_manual_quality_override": settings.time_series_manual_quality_override, "time_series_require_fresh_model": settings.time_series_require_fresh_model, diff --git a/crypto_spot_bot/storage.py b/crypto_spot_bot/storage.py index b5370df..03f45bb 100644 --- a/crypto_spot_bot/storage.py +++ b/crypto_spot_bot/storage.py @@ -11,8 +11,14 @@ from typing import Any, Iterator from crypto_spot_bot.models import Position, Signal, Trade, utc_now -MAX_SIGNAL_DIAGNOSTICS_BYTES = 16 * 1024 -PRUNE_BATCH_SIZE = 1000 +MAX_SIGNAL_DIAGNOSTICS_BYTES = 4 * 1024 +PRUNE_BATCH_SIZE = 5000 +MAX_RUNTIME_ROWS = { + "signals": 50_000, + "equity": 100_000, + "events": 20_000, + "llm_advice": 20_000, +} _STORED_FORECAST_KEYS = { "enabled", "usable", @@ -65,6 +71,9 @@ class Storage: def init_schema(self) -> None: with self.connect() as conn: + # New runtime databases reclaim deleted telemetry pages incrementally. + # Existing databases keep their current mode until compacted once. + conn.execute("PRAGMA auto_vacuum=INCREMENTAL") conn.execute("PRAGMA journal_mode=WAL") conn.executescript( """ @@ -628,21 +637,41 @@ class Storage: deleted: dict[str, int] = {} for table in ("signals", "equity", "events", "llm_advice"): with self.connect() as conn: - # Keep write locks short on large runtime databases. Each maintenance - # cycle removes at most one bounded batch per table. - cursor = conn.execute( - f""" - DELETE FROM {table} - WHERE id IN ( - SELECT id FROM {table} - WHERE created_at < ? - ORDER BY id - LIMIT ? + 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 + cap_boundary = max(0, max_id - MAX_RUNTIME_ROWS[table]) + removed = 0 + if cap_boundary > 0: + cursor = conn.execute( + f""" + DELETE FROM {table} + WHERE id IN ( + SELECT id FROM {table} + WHERE id <= ? + ORDER BY id + LIMIT ? + ) + """, + (cap_boundary, PRUNE_BATCH_SIZE), ) - """, - (cutoff, PRUNE_BATCH_SIZE), - ) - deleted[table] = max(0, int(cursor.rowcount)) + removed = max(0, int(cursor.rowcount)) + remaining = max(0, PRUNE_BATCH_SIZE - removed) + if remaining: + cursor = conn.execute( + f""" + DELETE FROM {table} + WHERE id IN ( + SELECT id FROM {table} + WHERE created_at < ? + ORDER BY id + LIMIT ? + ) + """, + (cutoff, remaining), + ) + removed += max(0, int(cursor.rowcount)) + deleted[table] = removed + conn.execute("PRAGMA incremental_vacuum(512)") return deleted def clear_all(self) -> None: diff --git a/crypto_spot_bot/strategy.py b/crypto_spot_bot/strategy.py index 6a87bb3..e8afa79 100644 --- a/crypto_spot_bot/strategy.py +++ b/crypto_spot_bot/strategy.py @@ -25,6 +25,35 @@ class SpotStrategy: trend_candles: list[Candle] | None = None, ) -> Signal: if self.settings.strategy_mode == "torch_forecast": + fallback_reasons = torch_model_readiness_reasons(self.settings, forecast or {}) + if self.settings.time_series_trend_fallback_enabled and fallback_reasons: + fallback = _trend_macd_entry_signal( + settings=self.settings, + symbol=symbol, + candles=candles, + trend_candles=trend_candles or [], + ticker=ticker, + open_positions_for_symbol=open_positions_for_symbol, + account=account, + ) + diagnostics = dict(fallback.diagnostics) + diagnostics.update( + { + "strategy_mode": "torch_forecast", + "trade_mode": "TREND_MACD_FALLBACK", + "entry_path": "trend_macd_fallback", + "forecast_fallback_active": True, + "forecast_fallback_reasons": fallback_reasons, + "forecast": forecast or {}, + } + ) + return Signal( + fallback.symbol, + fallback.action, + fallback.confidence, + f"torch_forecast fallback: {fallback.reason}", + diagnostics, + ) return _torch_forecast_entry_signal( settings=self.settings, symbol=symbol, @@ -368,6 +397,24 @@ class SpotStrategy: forecast: dict | None = None, ) -> Signal: if self.settings.strategy_mode == "torch_forecast": + if str(position.entry_diagnostics.get("entry_path", "")) == "trend_macd_fallback": + fallback = _trend_macd_exit_signal(self.settings, position, candles, ticker) + diagnostics = dict(fallback.diagnostics) + diagnostics.update( + { + "strategy_mode": "torch_forecast", + "trade_mode": "TREND_MACD_FALLBACK", + "entry_path": "trend_macd_fallback", + "forecast_fallback_active": True, + } + ) + return Signal( + fallback.symbol, + fallback.action, + fallback.confidence, + f"torch_forecast fallback: {fallback.reason}", + diagnostics, + ) return _torch_forecast_exit_signal(self.settings, position, candles, ticker, forecast or {}) if self.settings.strategy_mode == "trend_macd": return _trend_macd_exit_signal(self.settings, position, candles, ticker) @@ -1053,6 +1100,21 @@ def _is_torch_forecast(forecast: dict) -> bool: return bool(forecast.get("usable", False)) and model in {"torch_lstm", "torch_gru"} +def torch_model_readiness_reasons(settings: Settings, forecast: dict) -> list[str]: + reasons: list[str] = [] + if not _is_torch_forecast(forecast): + reasons.append("torch_model_unavailable") + if ( + settings.time_series_require_quality_gate + and not settings.time_series_manual_quality_override + and forecast.get("quality_gate_passed") is not True + ): + reasons.append("quality_gate_not_passed") + if settings.time_series_require_fresh_model and forecast.get("model_fresh") is not True: + reasons.append("model_not_fresh") + return reasons + + def _missing_torch_model(forecast: dict) -> bool: model = str(forecast.get("model", "")).strip().lower() reason = str(forecast.get("reason", "")).lower() diff --git a/crypto_spot_bot/time_series.py b/crypto_spot_bot/time_series.py index 43e9e9b..07e32a5 100644 --- a/crypto_spot_bot/time_series.py +++ b/crypto_spot_bot/time_series.py @@ -213,6 +213,7 @@ class TimeSeriesForecaster: probability=self.settings.time_series_min_probability_up, confidence=self.settings.time_series_min_confidence, ) + symbol_eligible = _calibration_symbol_eligible(calibration, symbol) entry = _torch_recurrent_entry(symbol, artifact) model = _torch_recurrent_model_name(symbol, artifact) clip = _clamp(_float_entry(entry or {}, "clip", 8.0), 1.0, 50.0) @@ -274,7 +275,8 @@ class TimeSeriesForecaster: ) conservative_return_percent = min(expected_return_percent, q50_percent) block_entry = bool( - (expected_return_percent <= -min_edge and probability_up <= 0.45) + not symbol_eligible + or (expected_return_percent <= -min_edge and probability_up <= 0.45) or (q50_percent <= -min_edge and probability_up <= 0.48) ) reason = _reason( @@ -284,6 +286,8 @@ class TimeSeriesForecaster: skill=skill, block_entry=block_entry, ) + if not symbol_eligible: + reason = "symbol excluded by train-only calibration" return TimeSeriesForecast( enabled=True, usable=True, @@ -344,7 +348,9 @@ class TimeSeriesForecaster: min_edge=min_edge, max_adjustment=self.settings.time_series_max_adjustment, ) - block_entry = bool(expected_return_percent <= -min_edge and probability_up <= 0.45) + block_entry = bool( + not symbol_eligible or (expected_return_percent <= -min_edge and probability_up <= 0.45) + ) reason = _reason( model=model, expected_return_percent=expected_return_percent, @@ -352,6 +358,8 @@ class TimeSeriesForecaster: skill=skill, block_entry=block_entry, ) + if not symbol_eligible: + reason = "symbol excluded by train-only calibration" return TimeSeriesForecast( enabled=True, usable=True, @@ -496,6 +504,16 @@ def _calibrated_thresholds( } +def _calibration_symbol_eligible(calibration: dict[str, Any], symbol: str | None) -> bool: + if not isinstance(calibration, dict) or "eligible_symbols" not in calibration: + return True + eligible = calibration.get("eligible_symbols") + if not isinstance(eligible, list) or not symbol: + return False + allowed = {str(value).strip().upper() for value in eligible if str(value).strip()} + return symbol.strip().upper() in allowed + + def _model_freshness(artifact: dict[str, Any], max_age_hours: float) -> tuple[str, float | None, bool]: raw = str(artifact.get("created_at", "")).strip() if isinstance(artifact, dict) else "" if not raw: @@ -998,7 +1016,12 @@ def _torch_recurrent_entry(symbol: str | None, artifact: dict[str, Any]) -> dict entry = default if isinstance(default, dict) else None if not isinstance(entry, dict): return None - if not isinstance(entry.get("state_dict"), dict): + members = entry.get("ensemble_members") + has_member_state = isinstance(members, list) and any( + isinstance(member, dict) and isinstance(member.get("state_dict"), dict) + for member in members + ) + if not isinstance(entry.get("state_dict"), dict) and not has_member_state: return None return entry diff --git a/crypto_spot_bot/training_coordination.py b/crypto_spot_bot/training_coordination.py index 4969020..3c88724 100644 --- a/crypto_spot_bot/training_coordination.py +++ b/crypto_spot_bot/training_coordination.py @@ -23,7 +23,9 @@ ALLOWED_TRAINING_ARTIFACTS = { RUNNING_TIMEOUT = timedelta(hours=12) ONLINE_WINDOW = timedelta(minutes=3) MAX_ARTIFACT_CHUNK_BYTES = 1024 * 1024 -MAX_ARTIFACT_BYTES = 64 * 1024 * 1024 +# Independent per-symbol ensembles are intentionally larger than pooled models. +# Keep a bounded limit, but leave enough room for the supported 12-symbol bundle. +MAX_ARTIFACT_BYTES = 256 * 1024 * 1024 MAX_ARTIFACT_CHUNKS = 1024 REQUIRED_MODEL_BUNDLE = set(ALLOWED_TRAINING_ARTIFACTS) @@ -378,13 +380,19 @@ def _safe_parameters(value: Any) -> dict[str, Any]: "dropouts", "epochs", "holdout_window", + "ensemble_seeds", + "selection_folds", + "learning_rate", + "weight_decay", + "pooled", "resume_candidate", } result = {key: value[key] for key in allowed if key in value} for key, low, high in ( - ("limit", 500, 5000), + ("limit", 500, 20000), ("epochs", 1, 200), ("holdout_window", 64, 1000), + ("selection_folds", 1, 12), ): if key not in result: continue @@ -406,9 +414,21 @@ def _safe_parameters(value: Any) -> dict[str, Any]: if item.strip().lower() in {"lstm", "gru"} ] result["architectures"] = ",".join(architectures) or "lstm,gru" - for key in ("lookbacks", "hidden_sizes", "layers", "dropouts"): + for key in ("lookbacks", "hidden_sizes", "layers", "dropouts", "ensemble_seeds"): if key in result: result[key] = str(result[key])[:200] + for key, low, high in ( + ("learning_rate", 0.00001, 0.1), + ("weight_decay", 0.0, 0.1), + ): + if key not in result: + continue + try: + result[key] = max(low, min(high, float(result[key]))) + except (TypeError, ValueError): + result.pop(key, None) + if "pooled" in result: + result["pooled"] = result["pooled"] is True if "resume_candidate" in result: result["resume_candidate"] = result["resume_candidate"] is True return result diff --git a/tests/conftest.py b/tests/conftest.py index a6feb13..d1b4d65 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -89,6 +89,7 @@ def make_settings(): time_series_probe_min_probability_up=0.55, time_series_probe_size_multiplier=0.40, time_series_rebound_fallback_enabled=True, + time_series_trend_fallback_enabled=False, stop_loss_percent=0.02, stop_loss_exit_enabled=True, take_profit_percent=0.035, diff --git a/tests/test_bybit.py b/tests/test_bybit.py index 0f2f8a8..ecdbdfe 100644 --- a/tests/test_bybit.py +++ b/tests/test_bybit.py @@ -1,5 +1,7 @@ from __future__ import annotations +import requests + from crypto_spot_bot.bybit import BybitClient, websocket_subscribe_message, _looks_like_leveraged_token, _looks_like_stablecoin @@ -86,6 +88,38 @@ def test_private_get_signs_the_same_query_it_sends(make_settings, tmp_path) -> N assert captured["headers"]["X-BAPI-SIGN"] +def test_public_get_recreates_failed_tls_session_before_retry(make_settings, tmp_path, monkeypatch) -> None: + client = BybitClient(make_settings(tmp_path)) + + class FailedSession: + def get(self, *_args, **_kwargs): + raise requests.exceptions.SSLError("invalid session id") + + class Response: + def raise_for_status(self): + return None + + def json(self): + return {"retCode": 0, "result": {"ok": True}} + + class WorkingSession: + def get(self, *_args, **_kwargs): + return Response() + + resets = [] + client.session = FailedSession() + + def reset_session() -> None: + resets.append(True) + client.session = WorkingSession() + + monkeypatch.setattr(client, "_reset_session", reset_session) + monkeypatch.setattr("crypto_spot_bot.bybit.time.sleep", lambda _seconds: None) + + assert client.public_get("/v5/market/kline", {"symbol": "BTCUSDT"}) == {"ok": True} + assert resets == [True] + + def test_websocket_subscribe_uses_configured_kline_interval() -> None: payload = websocket_subscribe_message(["BTCUSDT"], interval="60") diff --git a/tests/test_calibrate_thresholds.py b/tests/test_calibrate_thresholds.py index 0583ff8..5623c8b 100644 --- a/tests/test_calibrate_thresholds.py +++ b/tests/test_calibrate_thresholds.py @@ -1,13 +1,18 @@ from __future__ import annotations +from types import SimpleNamespace + from tools.calibrate_torch_thresholds import ( CalibrationResult, ForecastRecord, + _average_selected_predictions, _apply_platt_calibration, _choose_recommendation, + _full_backtest, _fit_platt_calibration, _entry_validation_skill, ) +from tools.train_torch_recurrent_forecaster import _ensemble_candidate def _result(*, trades: int, average: float, total: float, profit_factor: float) -> CalibrationResult: @@ -85,3 +90,81 @@ def test_entry_quality_never_falls_back_to_holdout_skill() -> None: assert _entry_validation_skill(entry) == 0.12 assert _entry_validation_skill({"skill": 0.99, "holdout_skill": 0.99}) == 0.0 + + +def test_batched_ensemble_averages_decoded_predictions() -> None: + averaged = _average_selected_predictions( + [ + {"expected_return": 0.01, "q50": 0.02, "probability_up": 0.6}, + {"expected_return": 0.03, "q50": 0.04, "probability_up": 0.8}, + ] + ) + + assert averaged == { + "expected_return": 0.02, + "q50": 0.03, + "probability_up": 0.7, + } + + +def test_multi_seed_export_does_not_duplicate_first_member_weights() -> None: + members = [ + { + "validation_mae": 0.1, + "state_dict": {"weight": [seed]}, + "head_weight": [[seed]], + "head_bias": [seed], + } + for seed in (7, 19) + ] + + exported = _ensemble_candidate(members, [7, 19]) + + assert exported["ensemble_size"] == 2 + assert exported["ensemble_seeds"] == [7, 19] + assert len(exported["ensemble_members"]) == 2 + assert "state_dict" not in exported + assert "head_weight" not in exported + + +def test_single_seed_export_keeps_only_top_level_weights() -> None: + exported = _ensemble_candidate( + [ + { + "validation_mae": 0.1, + "state_dict": {"weight": [7]}, + "head_weight": [[7]], + "head_bias": [7], + } + ], + [7], + ) + + assert exported["ensemble_size"] == 1 + assert exported["state_dict"] == {"weight": [7]} + assert "ensemble_members" not in exported + + +def test_full_backtest_never_uses_global_threshold_for_ineligible_symbol() -> None: + btc = [_record(index, 0.8, 1.0) for index in range(3)] + eth = [_record(index, 0.8, 1.0) for index in range(3)] + for record in eth: + record.symbol = "ETHUSDT" + thresholds = _result(trades=3, average=1.0, total=3.0, profit_factor=999.0) + + replay = _full_backtest( + btc + eth, + thresholds, + horizon=3, + round_trip_cost=0.0, + settings=SimpleNamespace( + stop_loss_percent=0.04, + take_profit_percent=0.035, + stop_loss_exit_enabled=True, + atr_trailing_multiplier=2.2, + ), + symbol_thresholds={"BTCUSDT": thresholds}, + require_symbol_thresholds=True, + ) + + assert {row["symbol"] for row in replay["symbol_breakdown"]} == {"BTCUSDT"} diff --git a/tests/test_storage.py b/tests/test_storage.py index d22200c..c323797 100644 --- a/tests/test_storage.py +++ b/tests/test_storage.py @@ -2,9 +2,11 @@ from __future__ import annotations import json from datetime import timedelta +from pathlib import Path from crypto_spot_bot.models import Signal, utc_now from crypto_spot_bot.storage import MAX_SIGNAL_DIAGNOSTICS_BYTES, PRUNE_BATCH_SIZE, Storage +from tools.compact_runtime_db import compact_database def test_hold_sampling_is_independent_for_each_reason_and_diagnostics_are_bounded(tmp_path) -> None: @@ -58,3 +60,29 @@ def test_prune_deletes_only_one_bounded_batch_per_table(tmp_path) -> None: assert deleted["signals"] == PRUNE_BATCH_SIZE assert len(storage.recent_signals(PRUNE_BATCH_SIZE + 10)) == 5 + + +def test_runtime_compaction_preserves_durable_state_and_bounds_telemetry(tmp_path) -> None: + database = tmp_path / "tradebot.sqlite3" + storage = Storage(database) + for index in range(10): + storage.insert_signal( + Signal("BTCUSDT", "BUY", 0.8, f"signal-{index}"), + hold_sample_seconds=0, + ) + storage.set_runtime("active", {"value": 1}) + + result = compact_database( + database, + recent_rows={"signals": 3, "equity": 0, "events": 0, "llm_advice": 0}, + ) + + compacted = Storage(database) + assert [row["reason"] for row in compacted.recent_signals(10)] == [ + "signal-9", + "signal-8", + "signal-7", + ] + assert compacted.get_runtime("active") == {"value": 1} + assert Path(result["backup"]).is_file() + assert result["rows"]["signals"] == 3 diff --git a/tests/test_strategy.py b/tests/test_strategy.py index 03ed8c4..2ce3a2b 100644 --- a/tests/test_strategy.py +++ b/tests/test_strategy.py @@ -566,6 +566,71 @@ def test_torch_forecast_blocks_failed_quality_gate(make_settings, tmp_path) -> N assert signal.diagnostics["checks"]["quality_gate_ok"] is False +def test_torch_forecast_uses_trend_fallback_when_model_is_not_ready(make_settings, tmp_path) -> None: + settings = make_settings( + tmp_path, + strategy_mode="torch_forecast", + time_series_trend_fallback_enabled=True, + time_series_require_quality_gate=True, + time_series_require_fresh_model=True, + max_position_usdt=50, + ) + strategy = SpotStrategy(settings) + ticker = Ticker("BTCUSDT", 105, 104.99, 105.01, 10_000_000, 1000, 1.0) + + signal = strategy.entry_signal( + "BTCUSDT", + _trend_entry_candles(), + ticker, + open_positions_for_symbol=0, + forecast={"usable": False, "model": "none", "quality_gate_passed": False}, + account={"equity": 100.0}, + trend_candles=_daily_trend_candles(), + ) + + assert signal.action == "BUY" + assert signal.diagnostics["trade_mode"] == "TREND_MACD_FALLBACK" + assert signal.diagnostics["entry_path"] == "trend_macd_fallback" + assert signal.diagnostics["forecast_fallback_reasons"] == [ + "torch_model_unavailable", + "quality_gate_not_passed", + "model_not_fresh", + ] + + +def test_torch_forecast_uses_trend_exit_for_fallback_position(make_settings, tmp_path) -> None: + settings = make_settings( + tmp_path, + strategy_mode="torch_forecast", + time_series_trend_fallback_enabled=True, + ) + strategy = SpotStrategy(settings) + candles = _trend_entry_candles() + candles[-2].macd = 0.2 + candles[-2].macd_signal = 0.0 + candles[-1].macd = -0.1 + candles[-1].macd_signal = 0.0 + position = Position( + 1, + "BTCUSDT", + 1, + 100, + 100, + 0.1, + 96, + 120, + 100, + entry_diagnostics={"entry_path": "trend_macd_fallback"}, + ) + ticker = Ticker("BTCUSDT", 104, 103.99, 104.01, 1_000_000, 100, 0) + + signal = strategy.exit_signal(position, candles, ticker, forecast={}) + + assert signal.action == "SELL" + assert signal.diagnostics["trade_mode"] == "TREND_MACD_FALLBACK" + assert "MACD" in signal.reason + + def test_torch_forecast_allows_explicit_manual_quality_override(make_settings, tmp_path) -> None: settings = make_settings( tmp_path, diff --git a/tests/test_time_series.py b/tests/test_time_series.py index 10a1657..116d041 100644 --- a/tests/test_time_series.py +++ b/tests/test_time_series.py @@ -334,6 +334,29 @@ def test_time_series_forecaster_uses_symbol_calibration(make_settings, tmp_path) assert forecast.calibrated_min_confidence == 0.45 +def test_time_series_forecaster_blocks_symbol_outside_train_only_allowlist(make_settings, tmp_path) -> None: + artifact_path = tmp_path / "lstm_forecaster.json" + _write_torch_gru_artifact(artifact_path, head_bias=0.2) + (tmp_path / "torch_threshold_calibration.json").write_text( + json.dumps( + { + "validation": {"status": "pass", "passed": True}, + "eligible_symbols": ["ETHUSDT"], + } + ), + encoding="utf-8", + ) + settings = make_settings(tmp_path, time_series_lstm_model_path=artifact_path) + + forecast = TimeSeriesForecaster(settings).forecast( + _candles_from_returns([0.0001] * 140), symbol="BTCUSDT" + ) + + assert forecast.usable is True + assert forecast.block_entry is True + assert forecast.reason == "symbol excluded by train-only calibration" + + def test_time_series_forecaster_averages_ensemble_members(make_settings, tmp_path) -> None: artifact_path = tmp_path / "lstm_forecaster.json" _write_torch_gru_artifact(artifact_path, head_bias=0.9) @@ -343,6 +366,9 @@ def test_time_series_forecaster_averages_ensemble_members(make_settings, tmp_pat {"state_dict": entry["state_dict"], "head_weight": [0.0, 0.0], "head_bias": bias} for bias in (0.1, 0.3) ] + entry.pop("state_dict") + entry.pop("head_weight") + entry.pop("head_bias") artifact_path.write_text(json.dumps(artifact), encoding="utf-8") settings = make_settings( tmp_path, diff --git a/tests/test_training_coordination.py b/tests/test_training_coordination.py index 264c4f7..14bf8d4 100644 --- a/tests/test_training_coordination.py +++ b/tests/test_training_coordination.py @@ -48,6 +48,33 @@ def test_training_coordinator_preserves_boolean_resume_candidate_parameter(tmp_p assert requested["job"]["parameters"] == {"resume_candidate": True} +def test_training_coordinator_sanitizes_independent_training_parameters(tmp_path) -> None: + coordinator = TrainingCoordinator(tmp_path) + + requested = coordinator.request_retrain( + { + "source": "recovery", + "parameters": { + "pooled": False, + "limit": 6000, + "ensemble_seeds": "7,19", + "selection_folds": 3, + "learning_rate": 0.0007, + "weight_decay": 0.0005, + }, + } + ) + + assert requested["job"]["parameters"] == { + "pooled": False, + "limit": 6000, + "ensemble_seeds": "7,19", + "selection_folds": 3, + "learning_rate": 0.0007, + "weight_decay": 0.0005, + } + + def test_training_coordinator_reports_worker_identity_from_heartbeat(tmp_path) -> None: coordinator = TrainingCoordinator(tmp_path) diff --git a/tools/calibrate_torch_thresholds.py b/tools/calibrate_torch_thresholds.py index 62c45bf..ff676e7 100644 --- a/tools/calibrate_torch_thresholds.py +++ b/tools/calibrate_torch_thresholds.py @@ -159,6 +159,7 @@ def main() -> None: settings=settings, ) symbol_recommendations: dict[str, dict[str, Any]] = {} + symbol_threshold_results: dict[str, CalibrationResult] = {} for symbol in symbols: symbol_records = [record for record in records if record.symbol == symbol] symbol_results = _calibrate_strategy( @@ -177,7 +178,8 @@ def main() -> None: ) if symbol_results else None if symbol_selected is not None: symbol_recommendations[symbol] = _result_dict(symbol_selected) - calibration_insufficient = recommended is None + symbol_threshold_results[symbol] = symbol_selected + calibration_insufficient = recommended is None or not symbol_threshold_results if recommended is None: recommended = _empty_recommendation( _float_grid(args.edge_grid), @@ -185,6 +187,16 @@ def main() -> None: _float_grid(args.confidence_grid), ) full_backtest = {**_stats([]), "trades_detail": [], "symbol_breakdown": []} + elif symbol_threshold_results: + full_backtest = _full_backtest( + records, + recommended, + horizon=horizon, + round_trip_cost=round_trip_cost, + settings=settings, + symbol_thresholds=symbol_threshold_results, + require_symbol_thresholds=True, + ) print("\nRECOMMENDED") print(_result_line(recommended)) print("\nFULL_REPLAY") @@ -248,6 +260,7 @@ def main() -> None: "recommended": _result_dict(deployment_recommended), "calibration_insufficient": calibration_insufficient, "symbol_recommendations": deployment_symbol_recommendations, + "eligible_symbols": sorted(deployment_symbol_recommendations), "full_replay": full_backtest, "walk_forward": walk_forward, "benchmark": benchmark, @@ -419,8 +432,8 @@ def _batch_forecast_records( horizons = _entry_target_horizons(entry) if not horizons: return None - model = _build_torch_model(entry, model_name) - if model is None: + models = _build_torch_models(entry, model_name) + if not models: return None lookback = int(_clamp(_float_entry(entry, "lookback", 64.0), 4.0, 512.0)) @@ -438,7 +451,8 @@ def _batch_forecast_records( records: list[ForecastRecord] = [] skill = _entry_validation_skill(entry) - model.eval() + for model in models: + model.eval() with torch.no_grad(): for offset in range(0, len(indices), max(1, batch_size)): batch_indices = indices[offset : offset + max(1, batch_size)] @@ -453,17 +467,25 @@ def _batch_forecast_records( for index in batch_indices ] batch = torch.tensor(windows, dtype=torch.float32) - outputs = model(batch).detach().cpu().tolist() - for index, output in zip(batch_indices, outputs): - selected = _decode_selected_output( - output, - entry=entry, - candles=candles, - closes=closes, - index=index, - horizon=decision_horizon, - clip=clip, - round_trip_cost=round_trip_cost, + outputs_by_model = [model(batch).detach().cpu().tolist() for model in models] + for batch_offset, index in enumerate(batch_indices): + selected = _average_selected_predictions( + [ + decoded + for outputs in outputs_by_model + if ( + decoded := _decode_selected_output( + outputs[batch_offset], + entry=entry, + candles=candles, + closes=closes, + index=index, + horizon=decision_horizon, + clip=clip, + round_trip_cost=round_trip_cost, + ) + ) is not None + ] ) if selected is None: continue @@ -502,11 +524,21 @@ def _batch_forecast_records( return records +def _build_torch_models(entry: dict[str, Any], model_name: str) -> list[Any]: + members = entry.get("ensemble_members") + if isinstance(members, list) and members: + base = {key: value for key, value in entry.items() if key != "ensemble_members"} + models = [ + _build_torch_model({**base, **member}, model_name) + for member in members + if isinstance(member, dict) + ] + return [model for model in models if model is not None] + model = _build_torch_model(entry, model_name) + return [model] if model is not None else [] + + def _build_torch_model(entry: dict[str, Any], model_name: str) -> Any | None: - if isinstance(entry.get("ensemble_members"), list) and entry["ensemble_members"]: - # Ensemble inference is handled by the shared pure-Python runtime so - # calibration and production use the exact same averaging path. - return None if torch is None or RecurrentReturnModel is None: return None architecture = "lstm" if model_name == "torch_lstm" else "gru" if model_name == "torch_gru" else "" @@ -560,6 +592,15 @@ def _build_torch_model(entry: dict[str, Any], model_name: str) -> Any | None: return model +def _average_selected_predictions(rows: list[dict[str, float]]) -> dict[str, float] | None: + if not rows: + return None + return { + name: sum(float(row[name]) for row in rows) / len(rows) + for name in ("expected_return", "q50", "probability_up") + } + + def _decode_selected_output( output: list[float], *, @@ -638,6 +679,7 @@ def _full_backtest( settings: Any, detail_limit: int = 50, symbol_thresholds: dict[str, CalibrationResult] | None = None, + require_symbol_thresholds: bool = False, ) -> dict[str, Any]: positions: dict[str, dict[str, Any]] = {} trades: list[float] = [] @@ -708,6 +750,8 @@ def _full_backtest( if record.symbol in positions: continue + if require_symbol_thresholds and record.symbol not in (symbol_thresholds or {}): + continue if _candidate_allows(record, active_thresholds.edge, active_thresholds.probability, active_thresholds.confidence): positions[record.symbol] = { "entry_price": record.next_open, @@ -907,6 +951,7 @@ def _walk_forward( settings=settings, detail_limit=0, symbol_thresholds=symbol_thresholds, + require_symbol_thresholds=True, ) test_rows = test_backtest.get("trades_detail", []) test_trades = [float(row.get("net_percent", 0.0) or 0.0) for row in test_rows if isinstance(row, dict)] @@ -921,6 +966,7 @@ def _walk_forward( "symbol_thresholds": { symbol: _result_dict(value) for symbol, value in symbol_thresholds.items() }, + "eligible_symbols": sorted(symbol_thresholds), "probability_calibration": probability_calibration, "test": {key: value for key, value in test_backtest.items() if key != "trades_detail"}, } diff --git a/tools/compact_runtime_db.py b/tools/compact_runtime_db.py new file mode 100644 index 0000000..d775af5 --- /dev/null +++ b/tools/compact_runtime_db.py @@ -0,0 +1,136 @@ +from __future__ import annotations + +import argparse +import json +import sqlite3 +import sys +from pathlib import Path +from typing import Any + +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.storage import Storage + + +PRESERVED_TABLES = ("positions", "trades", "runtime", "orders") +DEFAULT_RECENT_ROWS = { + "signals": 5_000, + "equity": 5_000, + "events": 2_000, + "llm_advice": 1_000, +} + + +def compact_database( + database: Path, + *, + recent_rows: dict[str, int] | None = None, + backup: Path | None = None, +) -> dict[str, Any]: + database = database.resolve() + if not database.is_file(): + raise FileNotFoundError(database) + limits = dict(DEFAULT_RECENT_ROWS) + if recent_rows: + limits.update({key: max(0, int(value)) for key, value in recent_rows.items()}) + temp = database.with_name(database.name + ".compact") + backup = (backup or database.with_name(database.name + ".precompact.bak")).resolve() + if temp.exists(): + temp.unlink() + if backup.exists(): + raise FileExistsError(f"backup already exists: {backup}") + + source_bytes = database.stat().st_size + Storage(temp) + counts: dict[str, int] = {} + conn = sqlite3.connect(temp) + try: + conn.execute("PRAGMA foreign_keys=OFF") + conn.execute("ATTACH DATABASE ? AS source", (str(database),)) + for table in PRESERVED_TABLES: + counts[table] = _copy_table(conn, table, limit=None) + for table, limit in limits.items(): + counts[table] = _copy_table(conn, table, limit=limit) + conn.commit() + # Check only the newly built main database. The attached multi-gigabyte + # source is preserved as the rollback copy and must not be rescanned here. + integrity = str(conn.execute("PRAGMA main.integrity_check").fetchone()[0]) + if integrity.lower() != "ok": + raise RuntimeError(f"compacted database integrity check failed: {integrity}") + conn.execute("DETACH DATABASE source") + conn.execute("PRAGMA wal_checkpoint(TRUNCATE)") + conn.execute("PRAGMA journal_mode=DELETE") + conn.commit() + finally: + conn.close() + + database.replace(backup) + temp.replace(database) + compacted_bytes = database.stat().st_size + return { + "database": str(database), + "backup": str(backup), + "source_bytes": source_bytes, + "compacted_bytes": compacted_bytes, + "reclaimed_bytes": max(0, source_bytes - compacted_bytes), + "rows": counts, + } + + +def _copy_table(conn: sqlite3.Connection, table: str, *, limit: int | None) -> int: + destination_columns = _columns(conn, "main", table) + source_columns = set(_columns(conn, "source", table)) + columns = [column for column in destination_columns if column in source_columns] + if not columns: + return 0 + quoted = ", ".join(f'"{column}"' for column in columns) + if limit is None: + conn.execute( + f'INSERT INTO main."{table}" ({quoted}) SELECT {quoted} FROM source."{table}"' + ) + elif limit > 0: + conn.execute( + f'INSERT INTO main."{table}" ({quoted}) ' + f'SELECT {quoted} FROM source."{table}" ORDER BY id DESC LIMIT ?', + (limit,), + ) + row = conn.execute(f'SELECT COUNT(*) FROM main."{table}"').fetchone() + return int(row[0] if row else 0) + + +def _columns(conn: sqlite3.Connection, schema: str, table: str) -> list[str]: + return [str(row[1]) for row in conn.execute(f'PRAGMA {schema}.table_info("{table}")')] + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Atomically compact the TradeBot runtime database while preserving durable trading state." + ) + parser.add_argument("--database", required=True) + parser.add_argument("--backup", default="") + parser.add_argument("--signals", type=int, default=DEFAULT_RECENT_ROWS["signals"]) + parser.add_argument("--equity", type=int, default=DEFAULT_RECENT_ROWS["equity"]) + parser.add_argument("--events", type=int, default=DEFAULT_RECENT_ROWS["events"]) + parser.add_argument("--llm-advice", type=int, default=DEFAULT_RECENT_ROWS["llm_advice"]) + return parser.parse_args() + + +def main() -> None: + args = _parse_args() + result = compact_database( + Path(args.database), + backup=Path(args.backup) if args.backup else None, + recent_rows={ + "signals": args.signals, + "equity": args.equity, + "events": args.events, + "llm_advice": args.llm_advice, + }, + ) + print(json.dumps(result, ensure_ascii=False, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/tools/run_torch_retrain.ps1 b/tools/run_torch_retrain.ps1 index 6f4fd53..bc17ca6 100644 --- a/tools/run_torch_retrain.ps1 +++ b/tools/run_torch_retrain.ps1 @@ -27,6 +27,7 @@ param( [string]$PiRoot = "", [string]$PiSshKeyPath = "", [switch]$NoPiRestart, + [switch]$Pooled, [switch]$SkipGuard, [switch]$ResumeCandidate ) @@ -148,13 +149,13 @@ if (-not $Horizons) { $Horizons = if ($env:TORCH_RETRAIN_HORIZONS) { $env:TORCH_ if (-not $Features -and $env:TORCH_RETRAIN_FEATURES) { $Features = $env:TORCH_RETRAIN_FEATURES } if (-not $ContextSymbols -and $env:TORCH_RETRAIN_CONTEXT_SYMBOLS) { $ContextSymbols = $env:TORCH_RETRAIN_CONTEXT_SYMBOLS } if ($Seed -le 0 -and $env:TORCH_RETRAIN_SEED) { $Seed = [int]$env:TORCH_RETRAIN_SEED } -if (-not $EnsembleSeeds) { $EnsembleSeeds = if ($env:TORCH_RETRAIN_ENSEMBLE_SEEDS) { $env:TORCH_RETRAIN_ENSEMBLE_SEEDS } else { "7,19,43" } } +if (-not $EnsembleSeeds) { $EnsembleSeeds = if ($env:TORCH_RETRAIN_ENSEMBLE_SEEDS) { $env:TORCH_RETRAIN_ENSEMBLE_SEEDS } else { "7,19" } } if ($SelectionFolds -le 0) { $SelectionFolds = if ($env:TORCH_RETRAIN_SELECTION_FOLDS) { [int]$env:TORCH_RETRAIN_SELECTION_FOLDS } else { 3 } } if ($LearningRate -le 0) { $LearningRate = if ($env:TORCH_RETRAIN_LEARNING_RATE) { [double]$env:TORCH_RETRAIN_LEARNING_RATE } else { 0.0007 } } if ($WeightDecay -le 0) { $WeightDecay = if ($env:TORCH_RETRAIN_WEIGHT_DECAY) { [double]$env:TORCH_RETRAIN_WEIGHT_DECAY } else { 0.0005 } } if ($Epochs -le 0) { $Epochs = if ($env:TORCH_RETRAIN_EPOCHS) { [int]$env:TORCH_RETRAIN_EPOCHS } else { 70 } } if ($Patience -le 0) { $Patience = if ($env:TORCH_RETRAIN_PATIENCE) { [int]$env:TORCH_RETRAIN_PATIENCE } else { 8 } } -if ($HoldoutWindow -le 0) { $HoldoutWindow = if ($env:TORCH_RETRAIN_HOLDOUT_WINDOW) { [int]$env:TORCH_RETRAIN_HOLDOUT_WINDOW } else { 240 } } +if ($HoldoutWindow -le 0) { $HoldoutWindow = if ($env:TORCH_RETRAIN_HOLDOUT_WINDOW) { [int]$env:TORCH_RETRAIN_HOLDOUT_WINDOW } else { 1000 } } 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" } @@ -196,6 +197,12 @@ try { "--weight-decay", $WeightDecay.ToString([Globalization.CultureInfo]::InvariantCulture), "--output", $CandidateFile ) + if ($Pooled) { + $trainerArgs += "--pooled" + } + else { + $trainerArgs += "--no-pooled" + } if ($Symbols) { $trainerArgs += @("--symbols", $Symbols) } if ($Interval) { $trainerArgs += @("--interval", $Interval) } if ($EnvFile) { $trainerArgs += @("--env", $EnvFile) } @@ -236,7 +243,7 @@ try { "tools\calibrate_torch_thresholds.py", "--limit", $Limit.ToString(), "--calibration-window", ([Math]::Min(2400, [Math]::Max(1200, [int]($Limit / 2)))).ToString(), - "--min-trades", "60", + "--min-trades", "24", "--walk-forward-folds", "8", "--confidence-grid", "0.40" ) diff --git a/tools/train_torch_recurrent_forecaster.py b/tools/train_torch_recurrent_forecaster.py index 08f14d1..7ed072e 100644 --- a/tools/train_torch_recurrent_forecaster.py +++ b/tools/train_torch_recurrent_forecaster.py @@ -1020,12 +1020,22 @@ def _ensemble_candidate(members: list[dict[str, Any]], seeds: list[int]) -> dict "context_norm_weight", "context_norm_bias", ) - result["ensemble_members"] = [ - {name: member[name] for name in export_names if name in member} - | {"seed": seeds[index] if index < len(seeds) else index} - for index, member in enumerate(members) - ] result["ensemble_size"] = len(members) + result["ensemble_seeds"] = [seeds[index] if index < len(seeds) else index for index in range(len(members))] + if len(members) > 1: + result["ensemble_members"] = [ + {name: member[name] for name in export_names if name in member} + | {"seed": seeds[index] if index < len(seeds) else index} + for index, member in enumerate(members) + ] + # Ensemble inference uses the member payloads. Keeping the first + # member at the top level duplicated a complete network in every + # exported symbol and could push an otherwise valid artifact over + # the server upload limit. + for name in export_names: + result.pop(name, None) + else: + result.pop("ensemble_members", None) symbol_names = sorted( { symbol diff --git a/tools/windows_training_agent.py b/tools/windows_training_agent.py index a9f00b8..7fe5f8d 100644 --- a/tools/windows_training_agent.py +++ b/tools/windows_training_agent.py @@ -118,6 +118,8 @@ def run_retrain(args: argparse.Namespace, job_id: str, job: dict[str, Any], repo value = parameters.get(key) if value not in (None, ""): cmd.extend([ps_arg, str(value)]) + if parameters.get("pooled") is True: + cmd.append("-Pooled") if parameters.get("resume_candidate") is True: cmd.append("-ResumeCandidate") log(log_path, "Running retrain: " + " ".join(quote_for_log(part) for part in cmd))