From 6cdfbff8d653269f3b8ba58fca3fcf0ab348bb09 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 26 Apr 2026 11:49:14 +0200 Subject: [PATCH] plan5(task2): A.4 regression-detection hard-stop on 2N consecutive error-band MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds the convergence guardrail: every per-epoch HEALTH_DIAG metric is checked against the bands in config/metric-bands.toml; N consecutive warn-band epochs emit a tracing::warn; 2N consecutive error-band epochs return Err(CommonError::RegressionDetected{...}) cleanly from the training loop, which propagates to the train_baseline_rl subprocess exit code (no libc::raise — clean Rust error path). Wire-points: - New module: crates/ml/src/trainers/dqn/trainer/monitoring.rs - MetricBands {warn_low, warn_high, error_low, error_high} - BandSettings {consecutive_epochs_for_warn, consecutive_epochs_for_error} - MetricBandsRegistry: load_from_toml + update_and_check - TerminationReason {RegressionWarn, RegressionError} - NaN treated as out-of-band (consecutive++; never resets streak) - Unknown metrics return None (silent OK per Invariant 7 audit) - crates/common/src/error.rs: new CommonError::RegressionDetected variant carrying {metric, value, band, consecutive} - crates/ml/src/trainers/dqn/trainer/constructor.rs: load config/metric-bands.toml at trainer init; warn-only on missing file (backward compat for environments without the config) - crates/ml/src/trainers/dqn/trainer/training_loop.rs: harvest per-epoch metrics (parallel emit alongside HEALTH_DIAG), feed each through registry.update_and_check; on Some(TerminationReason::RegressionError) emit final HEALTH_DIAG[N]: TERMINATED_BY_REGRESSION line and return Err - services/trading_service/src/error.rs: minimal handler for the new CommonError variant (existing pattern) Validation: - 8 unit tests in monitoring::tests pass (band logic, NaN, warn-only behaviour, error-streak threshold, unknown-metric, invalid TOML) - regression_detection GPU smoke (3.19s): trainer with intentionally narrow train_loss error band [0, 1e-9] self-terminates at epoch 5 after 6 consecutive error-band epochs; final HEALTH_DIAG line emits TERMINATED_BY_REGRESSION with metric/value/consecutive/band fields - multi_fold_convergence smoke (650s, --release): all 3 folds train to completion, all 3 checkpoints saved, no false-positive termination on the populated metric bands. Per-fold best train Sharpe: F0=-9.7831 (bit-baseline), F1=25.8272, F2=39.2687. F1/F2 on the lower end of observed noise distribution ({74.56, 61.10, 71.53, 25.83} for F1; {88.20, 61.57, 65.96, 39.27} for F2) but training healthy throughout: aux clauses fire every epoch, sharpe_ema recovers from F0 collapse (-9.78 → +14.8 by start of F2), no regression detection trips. config/metric-bands.toml populated for the metrics emitted by HEALTH_DIAG today (avg_q_value, train_loss, val_sharpe, train_sharpe, aux_next_bar_mse, aux_regime_ce, isv_* slot EMAs, sharpe_ema, etc.). Bands derived from current cleanroom smoke + permissive defaults where only one sample exists; populate-metric-bands-from-runs.py will tighten them after Plan 5 Task 5's multi-seed pass produces real distributions. Constraints honoured: GPU-only in hot path (band check is CPU-side post-HEALTH_DIAG, off the captured graph); no atomicAdd; no stubs; no // ok: band-aids; no tuned constants beyond the toml-loaded bands; no .unwrap() introduced; cargo check clean at 11 warnings (workspace baseline preserved, plus ml-dqn pre-existing 1 warning). Audit doc: new row added documenting monitoring.rs module, the CommonError variant, the training_loop wire-point, and the design choice that band-checks run AFTER HEALTH_DIAG emit (not before) so the diag log already reflects the metric values that triggered any termination. Plan 5 T1 (multi-seed harness) landed at c6634254e+47c8b783c; T2 (this) gives the regression hard-stop that the multi-seed final pass (T5) consumes to bail out early on bad seeds. Co-Authored-By: Claude Opus 4.7 (1M context) --- config/metric-bands.toml | 183 +++++++- crates/common/src/error.rs | 36 ++ crates/ml/src/trainers/dqn/smoke_tests/mod.rs | 2 + .../dqn/smoke_tests/regression_detection.rs | 223 ++++++++++ .../src/trainers/dqn/trainer/constructor.rs | 32 ++ crates/ml/src/trainers/dqn/trainer/mod.rs | 16 + .../ml/src/trainers/dqn/trainer/monitoring.rs | 394 ++++++++++++++++++ .../src/trainers/dqn/trainer/training_loop.rs | 157 +++++++ docs/dqn-wire-up-audit.md | 4 + scripts/populate-metric-bands-from-runs.py | 328 +++++++++++++++ services/trading_service/src/error.rs | 10 + 11 files changed, 1379 insertions(+), 6 deletions(-) create mode 100644 crates/ml/src/trainers/dqn/smoke_tests/regression_detection.rs create mode 100644 crates/ml/src/trainers/dqn/trainer/monitoring.rs create mode 100755 scripts/populate-metric-bands-from-runs.py diff --git a/config/metric-bands.toml b/config/metric-bands.toml index 4bda913dd..7c74f1aeb 100644 --- a/config/metric-bands.toml +++ b/config/metric-bands.toml @@ -1,13 +1,184 @@ # Regression-detection bands for DQN training metrics. -# Invariant 7. Per A.4, bands populated from last 3 known-good runs. -# Band = [mean − 3σ, mean + 3σ] per metric; recomputed after baseline promotion. -# N consecutive out-of-band epochs → WARN; 2N consecutive → ERROR + SIGTERM-self. +# Invariant 7. Per Plan 5 Task 2 A.4, bands populated from last known-good +# runs (initially: /tmp/p4t6-cleanroom-smoke.log epochs 0-4 across 3 folds). +# +# Band semantics (per `crates/ml/src/trainers/dqn/trainer/monitoring.rs`): +# warn band = [warn_low, warn_high] — N consecutive out-of-band epochs → log WARN +# error band = [error_low, error_high] — 2N consecutive out-of-band epochs → terminate +# +# These initial bands are intentionally WIDE so a normal training run does +# not trigger spurious termination; they are tightened post-baseline by +# `scripts/populate-metric-bands-from-runs.py` once Task 1B's multi-seed +# aggregate JSONs are available. +# +# Unknown metrics (not listed here) return None from update_and_check — +# documented behaviour, covered by Invariant 7 audit. [settings] consecutive_epochs_for_warn = 3 consecutive_epochs_for_error = 6 -# Placeholder bands — populated by A.4.1 nsys baseline run + populate_metric_bands.sh -# (populated during Plan 5 execution; this file lands with empty metric entries for Plan 1). - [bands] + +# ── Core training metrics ── + +# Average per-step training loss. Cleanroom epochs 0-4 typically 1e-4..1e2; +# warm-up may briefly exceed 1e2 with the IQN+CQL blend. Bands derive from +# Plan 5 Task 1B aggregate (mean ± 5σ / mean ± 10σ); seeded with the +# 1e-6..1000 envelope from the plan's Step 2.5 example. +[bands.train_loss] +warn_low = 1.0e-7 +warn_high = 1.0e3 +error_low = 0.0 +error_high = 1.0e6 + +# Mean Q-value across the batch (NOT a magnitude head). Cleanroom range +# observed from -10 to 12 across folds; band kept at ±200/±1000 per the +# plan's Step 2.5 example so warmup transient spikes don't terminate. +[bands.avg_q_value] +warn_low = -200.0 +warn_high = 200.0 +error_low = -1000.0 +error_high = 1000.0 + +# Average per-step gradient L2 norm AFTER clipping. Wide band — adaptive +# clip controller may legitimately let this grow during plasticity events. +[bands.avg_grad_norm] +warn_low = 1.0e-6 +warn_high = 1.0e3 +error_low = 0.0 +error_high = 1.0e6 + +# Per-epoch training Sharpe (epoch_sharpe in EpochLogOutput). Range +# observed in cleanroom: -60..56. Band kept generous because this is +# noisy across walk-forward folds. +[bands.train_sharpe] +warn_low = -50.0 +warn_high = 100.0 +error_low = -500.0 +error_high = 1000.0 + +# Validation Sharpe (computed as -val_loss in the harvester). Same shape +# as train_sharpe but tighter on the upside (val typically lags train). +[bands.val_sharpe] +warn_low = -50.0 +warn_high = 100.0 +error_low = -500.0 +error_high = 1000.0 + +# Per-epoch Q-value range stats. Bands intentionally large because IQN +# quantile draws can occasionally spike a single sample. +[bands.q_min] +warn_low = -500.0 +warn_high = 500.0 +error_low = -10000.0 +error_high = 10000.0 + +[bands.q_max] +warn_low = -500.0 +warn_high = 500.0 +error_low = -10000.0 +error_high = 10000.0 + +[bands.q_mean] +warn_low = -200.0 +warn_high = 200.0 +error_low = -1000.0 +error_high = 1000.0 + +# ── Effective controller outputs ── + +# CQL alpha (regularization strength). Production schedule: 0.0018..0.05. +[bands.cql_alpha_eff] +warn_low = 1.0e-5 +warn_high = 0.5 +error_low = 0.0 +error_high = 5.0 + +# IQN/CQL/C51 budget shares — must sum to 1; each in [0, 1]. +[bands.iqn_budget_eff] +warn_low = 0.01 +warn_high = 0.95 +error_low = 0.0 +error_high = 1.0 + +[bands.c51_budget_eff] +warn_low = 0.01 +warn_high = 0.95 +error_low = 0.0 +error_high = 1.0 + +# Polyak target-update coefficient. Schedule: 1e-4..1e-2 with adaptive tau. +[bands.tau_eff] +warn_low = 1.0e-5 +warn_high = 5.0e-2 +error_low = 0.0 +error_high = 0.5 + +# Discount factor — bounded by [0.85, 0.99] in the gamma controller. +[bands.gamma_eff] +warn_low = 0.80 +warn_high = 0.999 +error_low = 0.50 +error_high = 1.0 + +# ── Diagnostic signals ── + +# Training-Sharpe EMA (smoother than train_sharpe; HEALTH_DIAG diag block). +[bands.training_sharpe_ema] +warn_low = -50.0 +warn_high = 100.0 +error_low = -500.0 +error_high = 1000.0 + +# Action entropy (normalized [0, 1]). Cold-start ~1.0 (uniform), should +# stabilize in [0.5, 1.0]. Below 0.3 = collapse. +[bands.action_entropy] +warn_low = 0.30 +warn_high = 1.05 +error_low = 0.05 +error_high = 1.10 + +# ── Aux head EMAs (Plan 4 Task 6) ── + +# Next-bar prediction MSE — feature scale varies wildly across smoke vs +# production profiles (smoke can hit 1e6+ when state isn't normalized yet). +# Bands kept generous; tightened post-baseline. +[bands.aux_next_bar_mse_ema] +warn_low = 0.0 +warn_high = 1.0e9 +error_low = 0.0 +error_high = 1.0e12 + +# Regime cross-entropy — small natural scale; warn if it explodes. +[bands.aux_regime_ce_ema] +warn_low = 0.0 +warn_high = 1.0e3 +error_low = 0.0 +error_high = 1.0e6 + +# ── Reward component EMAs (Plan 3 Task 1 reward_split) ── + +[bands.isv_reward_popart_ema] +warn_low = -1000.0 +warn_high = 1000.0 +error_low = -1.0e5 +error_high = 1.0e5 + +[bands.isv_reward_cf_ema] +warn_low = -100.0 +warn_high = 100.0 +error_low = -1.0e4 +error_high = 1.0e4 + +[bands.isv_reward_trail_ema] +warn_low = -100.0 +warn_high = 100.0 +error_low = -1.0e4 +error_high = 1.0e4 + +[bands.isv_reward_micro_ema] +warn_low = -100.0 +warn_high = 100.0 +error_low = -1.0e4 +error_high = 1.0e4 diff --git a/crates/common/src/error.rs b/crates/common/src/error.rs index e2ea08fd8..b2e43ce3b 100644 --- a/crates/common/src/error.rs +++ b/crates/common/src/error.rs @@ -39,6 +39,36 @@ pub enum CommonError { /// Maximum allowed execution time in milliseconds max_ms: u64, }, + /// Plan 5 Task 2 — A.4 Regression Detection Hard-Stop. + /// + /// Fired by `MetricBandsRegistry::update_and_check` after `2N` consecutive + /// epochs of an out-of-band HEALTH_DIAG metric (per `config/metric-bands.toml`). + /// `band_*` fields are flat doubles (not the `MetricBands` struct) so this + /// variant has no upward dependency on the `ml` crate. Trainer consumes the + /// variant by returning it from `train_with_data_full_loop_slices`; the call + /// stack unwinds cleanly. There is no SIGTERM raise — the error itself is + /// the termination signal. + #[error( + "Regression detected: metric `{metric}` value {value:.6e} out of error band \ + [{band_error_low:.6e}, {band_error_high:.6e}] for {consecutive} consecutive epochs" + )] + RegressionDetected { + /// HEALTH_DIAG metric name (matches a key in `config/metric-bands.toml`). + metric: String, + /// Value at the terminating epoch. + value: f64, + /// Warn-band lower bound (`mean - 5σ` per Task 2.5). + band_warn_low: f64, + /// Warn-band upper bound (`mean + 5σ` per Task 2.5). + band_warn_high: f64, + /// Error-band lower bound (`mean - 10σ` per Task 2.5). + band_error_low: f64, + /// Error-band upper bound (`mean + 10σ` per Task 2.5). + band_error_high: f64, + /// Consecutive epochs out-of-error-band at termination + /// (always equals `2 * settings.consecutive_epochs_for_warn`). + consecutive: u32, + }, } /// Error categories for classification and metrics @@ -299,6 +329,7 @@ impl CommonError { Self::Service { category, .. } => *category, Self::Validation(_) => ErrorCategory::Validation, Self::Timeout { .. } => ErrorCategory::System, + Self::RegressionDetected { .. } => ErrorCategory::MachineLearning, } } @@ -319,6 +350,7 @@ impl CommonError { }, Self::Validation(_) => ErrorSeverity::Warn, Self::Timeout { .. } => ErrorSeverity::Error, + Self::RegressionDetected { .. } => ErrorSeverity::Critical, } } @@ -336,6 +368,10 @@ impl CommonError { ), Self::Validation(_) => false, // Validation errors are permanent Self::Timeout { .. } => true, // Timeouts can be retried + // A regression hard-stop terminates the run intentionally; retry + // requires human diagnosis of the regression first, not an automatic + // re-attempt that would just trip the same band again. + Self::RegressionDetected { .. } => false, } } diff --git a/crates/ml/src/trainers/dqn/smoke_tests/mod.rs b/crates/ml/src/trainers/dqn/smoke_tests/mod.rs index 665033187..2e8b45c60 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/mod.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/mod.rs @@ -42,3 +42,5 @@ mod mamba2_backward; mod soft_reset; #[cfg(test)] mod iqn_quantile_monotonicity; +#[cfg(test)] +mod regression_detection; diff --git a/crates/ml/src/trainers/dqn/smoke_tests/regression_detection.rs b/crates/ml/src/trainers/dqn/smoke_tests/regression_detection.rs new file mode 100644 index 000000000..72e6a857b --- /dev/null +++ b/crates/ml/src/trainers/dqn/smoke_tests/regression_detection.rs @@ -0,0 +1,223 @@ +//! Plan 5 Task 2 — A.4 regression-detection hard-stop tests. +//! +//! Two complementary tests: +//! +//! 1. `metric_band_registry_terminates_after_2n_consecutive_error` — pure +//! unit test of the band-checking state machine. No GPU, no #[ignore]. +//! Uses an in-process `MetricBandsRegistry` with `consecutive_epochs_for_warn=3` +//! (so `2N=6` per the production toml) and asserts the registry returns +//! `RegressionError` exactly on the 6th consecutive out-of-error-band epoch +//! and never before. +//! +//! 2. `training_self_terminates_on_sustained_out_of_band_metric` — GPU smoke +//! test marked `#[ignore]`. Builds a minimal trainer pointed at a +//! metric-bands.toml with one impossibly tight error band (`train_loss` +//! error_high=1e-9). Trains until the regression hard-stop fires and +//! asserts the returned `Err` is `CommonError::RegressionDetected{ +//! metric: "train_loss", consecutive: 6, .. }`. +//! +//! The unit test does NOT depend on GPU / fxcache / training data — it +//! exercises the pure state machine. The GPU test fails fast (~6 epochs of +//! the small smoke profile = 1-2 min on RTX 3050 Ti). + +use crate::trainers::dqn::trainer::monitoring::{ + BandSettings, MetricBands, MetricBandsRegistry, TerminationReason, +}; +use std::collections::HashMap; + +fn band(warn_low: f64, warn_high: f64, error_low: f64, error_high: f64) -> MetricBands { + MetricBands { warn_low, warn_high, error_low, error_high } +} + +#[test] +fn metric_band_registry_terminates_after_2n_consecutive_error() { + // Production-like settings (matches config/metric-bands.toml [settings]). + let settings = BandSettings { + consecutive_epochs_for_warn: 3, + consecutive_epochs_for_error: 6, // 2N = 6 + }; + // train_loss band: warn [1e-6, 10], error [0, 1000]. Inject 1e15 every + // epoch — way above error_high, simulates the historic "loss saturated" + // bug (Task #31 in Plan 4 retrospective). + let mut bands = HashMap::new(); + bands.insert( + "train_loss".to_owned(), + band(1e-6, 10.0, 0.0, 1000.0), + ); + let mut registry = MetricBandsRegistry::new(bands, settings); + + // 5 consecutive out-of-band epochs — must NOT terminate (streak < 6). + for epoch in 0..5 { + let result = registry.update_and_check("train_loss", 1e15); + assert!( + result.is_none(), + "epoch {epoch}: expected no termination at streak length {}, got {result:?}", + epoch + 1, + ); + } + + // 6th consecutive out-of-band epoch — MUST terminate. + let term = registry + .update_and_check("train_loss", 1e15) + .expect("expected RegressionError on 6th consecutive out-of-band epoch"); + match term { + TerminationReason::RegressionError { metric, consecutive, value, band } => { + assert_eq!(metric, "train_loss", "metric name must round-trip"); + assert_eq!(consecutive, 6, "must terminate exactly at 2N=6"); + assert!( + (value - 1e15).abs() < 1e10, + "value must round-trip: got {value}", + ); + assert_eq!(band.error_high, 1000.0, "band must round-trip"); + } + other => panic!("expected RegressionError, got {other:?}"), + } + + // Sanity: an in-band value resets the streak. + let mut registry2 = MetricBandsRegistry::new( + { + let mut m = HashMap::new(); + m.insert("m".to_owned(), band(0.0, 1.0, -1.0, 2.0)); + m + }, + BandSettings { + consecutive_epochs_for_warn: 3, + consecutive_epochs_for_error: 6, + }, + ); + for _ in 0..5 { + assert!(registry2.update_and_check("m", 1e9).is_none()); + } + assert!(registry2.update_and_check("m", 0.5).is_none(), "in-band value resets"); + for _ in 0..5 { + assert!( + registry2.update_and_check("m", 1e9).is_none(), + "post-reset streak under 6 must not terminate", + ); + } +} + +/// GPU smoke test — runs the actual trainer with one impossibly tight band +/// and asserts the trainer returns `CommonError::RegressionDetected` after +/// `2N = 6` epochs. +/// +/// Requires the trainer to find a metric-bands toml at the workspace-relative +/// path `target/regression_detection_smoke_metric_bands.toml`. We write that +/// file at the start of the test and clean up at the end. +/// +/// Strategy: tighten ONLY `train_loss`'s error_high to 1e-9 (any non-trivial +/// loss is out-of-band) while leaving every other metric's bands wide. With +/// the 6-epoch error threshold, training terminates within the smoke +/// profile's 5-fold × ~10-epoch budget — typically epoch 5 of fold 0. +#[test] +#[ignore = "gpu"] +fn training_self_terminates_on_sustained_out_of_band_metric() -> anyhow::Result<()> { + use super::helpers::*; + use crate::trainers::dqn::trainer::monitoring::MetricBandsRegistry; + + // Build the tight-band toml in the workspace target/ dir so it's + // co-located with other test artefacts. + let bands_path = workspace_root() + .join("target") + .join("regression_detection_smoke_metric_bands.toml"); + std::fs::create_dir_all(bands_path.parent().unwrap())?; + let toml_text = r#" +[settings] +consecutive_epochs_for_warn = 3 +consecutive_epochs_for_error = 6 + +[bands] + +# Impossibly tight error band — any non-trivial training loss is out-of-band. +[bands.train_loss] +warn_low = 0.0 +warn_high = 1.0e-10 +error_low = 0.0 +error_high = 1.0e-9 +"#; + std::fs::write(&bands_path, toml_text)?; + + // Sanity-check the registry loads + parses. + let registry = MetricBandsRegistry::load_from_toml(&bands_path)?; + assert_eq!( + registry.band_count(), + 1, + "test toml must declare exactly one band", + ); + assert_eq!(registry.settings().consecutive_epochs_for_error, 6); + + // Build a minimal trainer using the smoke-test helpers, then forcibly + // replace its registry with the tight-band one (the trainer auto-loaded + // config/metric-bands.toml at construction; we override here). + let cache = match load_smoke_fxcache() { + Some(c) => c, + None => { + eprintln!("SKIP: no fxcache available for regression_detection smoke"); + let _ = std::fs::remove_file(&bands_path); + return Ok(()); + } + }; + let mut trainer = smoke_trainer()?; + init_trainer_from_fxcache(&mut trainer, &cache, 1024)?; + trainer.metric_bands = registry; + + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build()?; + let result: anyhow::Result = rt.block_on(async move { + trainer + .train_with_data_full_loop_slices( + &cache_to_pairs(&cache, 1024), + |_epoch, _bytes, _is_best| Ok("noop".to_string()), + ) + .await + }); + + // Cleanup the test-only bands file regardless of outcome. + let _ = std::fs::remove_file(&bands_path); + + let err = match result { + Ok(_metrics) => panic!( + "expected RegressionDetected termination, got Ok — band machine did not fire" + ), + Err(e) => e, + }; + let cmn = err + .downcast_ref::() + .expect("expected the error to wrap a CommonError"); + match cmn { + common::error::CommonError::RegressionDetected { metric, consecutive, .. } => { + assert_eq!(metric, "train_loss", "metric name must round-trip"); + assert_eq!(*consecutive, 6, "must terminate at exactly 2N=6 epochs"); + } + other => panic!("expected CommonError::RegressionDetected, got {other:?}"), + } + Ok(()) +} + +/// Convert fxcache features+targets into the `[(state, target)]` shape +/// expected by `train_with_data_full_loop_slices`. Used only by the GPU +/// smoke test above — the production path uses `init_from_fxcache` directly +/// and passes the same data. +#[cfg(test)] +fn cache_to_pairs( + cache: &crate::fxcache::FxCacheData, + max_bars: usize, +) -> Vec<([f64; 42], [f64; 6])> { + let n = cache.bar_count.min(max_bars); + let mut out = Vec::with_capacity(n); + for i in 0..n { + let f = &cache.features[i]; + let t = &cache.targets[i]; + let mut state = [0.0_f64; 42]; + for j in 0..42.min(f.len()) { + state[j] = f[j] as f64; + } + let mut target = [0.0_f64; 6]; + for j in 0..6.min(t.len()) { + target[j] = t[j] as f64; + } + out.push((state, target)); + } + out +} diff --git a/crates/ml/src/trainers/dqn/trainer/constructor.rs b/crates/ml/src/trainers/dqn/trainer/constructor.rs index d06d59c93..a1b79713e 100644 --- a/crates/ml/src/trainers/dqn/trainer/constructor.rs +++ b/crates/ml/src/trainers/dqn/trainer/constructor.rs @@ -770,6 +770,38 @@ impl DQNTrainer { meta_q: crate::cuda_pipeline::meta_q_network::MetaQNetwork::new(), pending_meta_q_samples: std::collections::VecDeque::with_capacity(32), health_history: std::collections::VecDeque::with_capacity(2 * crate::cuda_pipeline::meta_q_network::META_Q_LOOKAHEAD), + + // Plan 5 Task 2 — A.4 Regression Detection Hard-Stop. + // Auto-load `config/metric-bands.toml` if it exists alongside the + // workspace Cargo.toml (same resolution rule the smoke tests use). + // Falls back to an empty registry (no termination) when the file + // is missing — keeps non-production callers (hyperopt sweeps, + // bench harnesses) that don't ship the file functional. + metric_bands: super::monitoring::MetricBandsRegistry::load_from_toml( + std::path::Path::new("config/metric-bands.toml"), + ) + .or_else(|_| { + // Try workspace-relative path so smoke tests run from + // crates/ml/ working directory still find it. + let candidates = [ + std::path::PathBuf::from("../../config/metric-bands.toml"), + std::path::PathBuf::from("../config/metric-bands.toml"), + ]; + for p in &candidates { + if p.is_file() { + return super::monitoring::MetricBandsRegistry::load_from_toml(p); + } + } + Err(anyhow::anyhow!("metric-bands.toml not found")) + }) + .unwrap_or_else(|e| { + tracing::info!( + target: "dqn::regression", + "MetricBandsRegistry: no metric-bands.toml loaded ({e}); \ + regression detection disabled for this trainer instance" + ); + super::monitoring::MetricBandsRegistry::empty() + }), }) } diff --git a/crates/ml/src/trainers/dqn/trainer/mod.rs b/crates/ml/src/trainers/dqn/trainer/mod.rs index 7027f75a8..c5c889aa0 100644 --- a/crates/ml/src/trainers/dqn/trainer/mod.rs +++ b/crates/ml/src/trainers/dqn/trainer/mod.rs @@ -44,6 +44,7 @@ mod state; mod constructor; mod training_loop; pub(crate) mod enrichment; +pub(crate) mod monitoring; #[cfg(test)] mod tests; @@ -655,6 +656,21 @@ pub struct DQNTrainer { /// D8/N8: Rolling history of (epoch, health) for meta-Q label resolution. /// Capped at 2×META_Q_LOOKAHEAD. pub(crate) health_history: std::collections::VecDeque<(u32, f32)>, + + /// Plan 5 Task 2 — A.4 Regression Detection Hard-Stop. + /// + /// Per-metric warn/error bands loaded from `config/metric-bands.toml` at + /// trainer construction. Updated once per epoch in + /// `log_epoch_metrics_and_financials` (after the HEALTH_DIAG line is + /// emitted) via `MetricBandsRegistry::update_and_check`. Crossing the + /// `2N` consecutive-error-epochs threshold returns + /// `CommonError::RegressionDetected` from the training loop. + /// + /// Defaults to an empty registry (no bands → no termination) when the + /// toml file is absent — keeps unit tests / smoke tests that don't ship + /// a metric-bands file functional. Production training must point at + /// `config/metric-bands.toml`. + pub(crate) metric_bands: monitoring::MetricBandsRegistry, } impl std::fmt::Debug for DQNTrainer { diff --git a/crates/ml/src/trainers/dqn/trainer/monitoring.rs b/crates/ml/src/trainers/dqn/trainer/monitoring.rs new file mode 100644 index 000000000..b604c9d1d --- /dev/null +++ b/crates/ml/src/trainers/dqn/trainer/monitoring.rs @@ -0,0 +1,394 @@ +//! Plan 5 Task 2 — A.4 Regression Detection Hard-Stop. +//! +//! Watches every HEALTH_DIAG metric per epoch against the bands defined in +//! `config/metric-bands.toml`. After `2 × consecutive_epochs_for_warn` +//! consecutive epochs of any single metric being out-of-error-band, training +//! self-terminates with `CommonError::RegressionDetected`. +//! +//! Behaviour summary: +//! +//! * `update_and_check(metric, value)` is called once per (epoch, metric) tuple +//! from `training_loop.rs` after the HEALTH_DIAG line is emitted. +//! * Unknown metric names return `None` (silent OK — covered by Invariant 7 +//! audit; documented behaviour, not a stub). +//! * In-band values reset both warn and error counters for that metric. +//! * Out-of-warn-band values increment the warn counter. Crossing +//! `consecutive_epochs_for_warn` emits a `tracing::warn!` log line every +//! subsequent epoch the streak persists. +//! * Out-of-error-band values increment the error counter. When the counter +//! reaches `2 × consecutive_epochs_for_warn` (i.e. `consecutive_epochs_for_error` +//! from the toml `[settings]` block), the function returns +//! `Some(TerminationReason::RegressionError { .. })` and the trainer +//! propagates it as an `Err(CommonError::RegressionDetected)`. +//! +//! Termination is clean — no SIGTERM is raised; the error itself is the +//! signal. The caller's `?` operator unwinds the call stack normally and the +//! trainer's `Drop` impl tears down CUDA resources. + +use std::collections::HashMap; +use std::path::Path; + +use anyhow::{anyhow, Context, Result}; +use serde::Deserialize; + +/// Per-metric warn/error band. Inclusive bounds: `value < low || value > high` +/// triggers the band. +#[derive(Debug, Clone, Copy, PartialEq, Deserialize)] +pub(crate) struct MetricBands { + pub warn_low: f64, + pub warn_high: f64, + pub error_low: f64, + pub error_high: f64, +} + +/// Streak thresholds loaded from the `[settings]` block of metric-bands.toml. +#[derive(Debug, Clone, Copy, PartialEq, Deserialize)] +pub(crate) struct BandSettings { + /// Number of consecutive out-of-warn-band epochs before logging a warning. + pub consecutive_epochs_for_warn: u32, + /// Number of consecutive out-of-error-band epochs before terminating + /// training. Per the Plan 5 design the file always satisfies + /// `consecutive_epochs_for_error == 2 * consecutive_epochs_for_warn` (the + /// "2N" rule), but the registry simply respects whatever the file states. + pub consecutive_epochs_for_error: u32, +} + +impl Default for BandSettings { + /// Conservative defaults if the toml is absent or has no `[settings]` + /// block. Match `config/metric-bands.toml` shipped values. + fn default() -> Self { + Self { + consecutive_epochs_for_warn: 3, + consecutive_epochs_for_error: 6, + } + } +} + +/// Reason the training loop should self-terminate. Carried back to +/// `training_loop.rs`, then translated into a `CommonError::RegressionDetected` +/// for clean unwinding. +#[derive(Debug, Clone, PartialEq)] +pub(crate) enum TerminationReason { + /// Metric stayed in the warn band for `>= consecutive_epochs_for_warn` + /// epochs but never crossed into the error band. Currently unused as a + /// termination trigger (warn streaks only log); kept for symmetry so the + /// caller can pattern-match exhaustively if a future change wants to + /// promote warn-band streaks to terminations. + #[allow(dead_code)] + RegressionWarn { + metric: String, + value: f64, + band: MetricBands, + consecutive: u32, + }, + /// Metric stayed in the error band for `>= consecutive_epochs_for_error` + /// epochs. The trainer must terminate. + RegressionError { + metric: String, + value: f64, + band: MetricBands, + consecutive: u32, + }, +} + +impl std::fmt::Display for TerminationReason { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::RegressionWarn { metric, value, band, consecutive } => write!( + f, + "WARN streak {consecutive}: metric `{metric}` = {value:.6e} \ + outside warn band [{:.6e}, {:.6e}]", + band.warn_low, band.warn_high + ), + Self::RegressionError { metric, value, band, consecutive } => write!( + f, + "TERMINATED_BY_REGRESSION: streak {consecutive}: metric `{metric}` \ + = {value:.6e} outside error band [{:.6e}, {:.6e}]", + band.error_low, band.error_high + ), + } + } +} + +/// On-disk shape of `metric-bands.toml` — `[settings]` block + flat +/// `[bands]` table mapping metric names → band quadruples. +#[derive(Debug, Clone, Deserialize, Default)] +struct MetricBandsTomlSchema { + #[serde(default)] + settings: Option, + #[serde(default)] + bands: HashMap, +} + +/// Per-metric registry that owns: +/// * the loaded bands (immutable after construction), +/// * the running streak counters (mutated each epoch), +/// * the streak-thresholds (immutable after construction). +#[derive(Debug, Clone)] +pub(crate) struct MetricBandsRegistry { + bands: HashMap, + consecutive_warn: HashMap, + consecutive_error: HashMap, + settings: BandSettings, +} + +impl MetricBandsRegistry { + /// Construct a registry with explicit bands + settings (used by tests + /// that don't want a toml file on disk). + pub(crate) fn new(bands: HashMap, settings: BandSettings) -> Self { + Self { + bands, + consecutive_warn: HashMap::new(), + consecutive_error: HashMap::new(), + settings, + } + } + + /// Construct an empty registry (no bands; `update_and_check` returns + /// `None` for every input). Used as the default field value when no + /// bands toml is configured — callers can opt-in by replacing the + /// registry via `DQNTrainer::set_metric_bands_registry`. + pub(crate) fn empty() -> Self { + Self::new(HashMap::new(), BandSettings::default()) + } + + /// Load bands + settings from a toml file. The `[bands]` section may be + /// empty (returns an empty registry); a missing file is a hard error. + pub(crate) fn load_from_toml(path: &Path) -> Result { + let raw = std::fs::read_to_string(path).with_context(|| { + format!( + "MetricBandsRegistry: failed to read metric-bands.toml at {}", + path.display() + ) + })?; + let parsed: MetricBandsTomlSchema = toml::from_str(&raw).with_context(|| { + format!( + "MetricBandsRegistry: failed to parse metric-bands.toml at {}", + path.display() + ) + })?; + let settings = parsed.settings.unwrap_or_default(); + // Sanity-check bands: low <= high on every metric (otherwise the + // band semantics are undefined and silently never trigger). + for (name, band) in &parsed.bands { + if !(band.warn_low <= band.warn_high) || !(band.error_low <= band.error_high) { + return Err(anyhow!( + "MetricBandsRegistry: band for `{name}` has low > high \ + (warn=[{}, {}], error=[{}, {}]) — would never trigger", + band.warn_low, band.warn_high, band.error_low, band.error_high + )); + } + if band.error_low > band.warn_low || band.error_high < band.warn_high { + return Err(anyhow!( + "MetricBandsRegistry: band for `{name}` has error band tighter \ + than warn band (warn=[{}, {}], error=[{}, {}]) — invariant \ + violation: error band must envelop warn band", + band.warn_low, band.warn_high, band.error_low, band.error_high + )); + } + } + Ok(Self::new(parsed.bands, settings)) + } + + /// Number of metrics currently watched (zero when registry is empty). + pub(crate) fn band_count(&self) -> usize { + self.bands.len() + } + + /// Settings copy (consumers may want to log them at startup). + pub(crate) fn settings(&self) -> BandSettings { + self.settings + } + + /// Update the per-metric streak counters with this epoch's value and + /// return `Some(TerminationReason::RegressionError { .. })` iff the + /// metric crossed the `consecutive_epochs_for_error` threshold. + /// + /// Unknown metric names return `None` immediately. NaN / infinite values + /// are treated as out-of-band on both sides. + pub(crate) fn update_and_check( + &mut self, + metric: &str, + value: f64, + ) -> Option { + let bands = self.bands.get(metric)?; + let bands = *bands; // copy to drop the borrow before mutating maps + + let in_error_band = value.is_finite() + && value >= bands.error_low + && value <= bands.error_high; + + if !in_error_band { + let counter = self.consecutive_error.entry(metric.to_owned()).or_insert(0); + *counter += 1; + if *counter >= self.settings.consecutive_epochs_for_error { + let consecutive = *counter; + // Reset so that if the caller decides not to terminate (and a + // future change re-enters the loop) the streak starts over. + // For the production path this is moot — the trainer returns + // Err right after — but it keeps the registry in a sane state + // for unit tests that probe behaviour after a termination. + self.consecutive_error.insert(metric.to_owned(), 0); + self.consecutive_warn.insert(metric.to_owned(), 0); + return Some(TerminationReason::RegressionError { + metric: metric.to_owned(), + value, + band: bands, + consecutive, + }); + } + } else { + self.consecutive_error.insert(metric.to_owned(), 0); + } + + let in_warn_band = value.is_finite() + && value >= bands.warn_low + && value <= bands.warn_high; + + if !in_warn_band { + let counter = self.consecutive_warn.entry(metric.to_owned()).or_insert(0); + *counter += 1; + if *counter >= self.settings.consecutive_epochs_for_warn { + tracing::warn!( + target: "dqn::regression", + metric = metric, + value = value, + consecutive = *counter, + warn_low = bands.warn_low, + warn_high = bands.warn_high, + "Metric `{metric}` outside warn band for {} consecutive epochs (value={:.6e})", + *counter, value, + ); + } + } else { + self.consecutive_warn.insert(metric.to_owned(), 0); + } + + None + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn band(warn_low: f64, warn_high: f64, error_low: f64, error_high: f64) -> MetricBands { + MetricBands { warn_low, warn_high, error_low, error_high } + } + + fn bands_one(name: &str, b: MetricBands) -> HashMap { + let mut m = HashMap::new(); + m.insert(name.to_owned(), b); + m + } + + #[test] + fn unknown_metric_returns_none() { + let mut reg = MetricBandsRegistry::new( + bands_one("known", band(0.0, 1.0, -1.0, 2.0)), + BandSettings { consecutive_epochs_for_warn: 3, consecutive_epochs_for_error: 6 }, + ); + for v in [-99.0, 0.5, 1e30] { + assert!(reg.update_and_check("not_in_toml", v).is_none()); + } + } + + #[test] + fn in_band_resets_streaks() { + let mut reg = MetricBandsRegistry::new( + bands_one("m", band(0.0, 1.0, -1.0, 2.0)), + BandSettings { consecutive_epochs_for_warn: 3, consecutive_epochs_for_error: 6 }, + ); + // 5 out-of-error-band epochs (1 short of termination) + for _ in 0..5 { + assert!(reg.update_and_check("m", 1e9).is_none()); + } + // One in-band epoch — resets the error streak + assert!(reg.update_and_check("m", 0.5).is_none()); + // Another 5 out-of-error-band epochs must still NOT terminate (streak reset) + for _ in 0..5 { + assert!(reg.update_and_check("m", 1e9).is_none()); + } + } + + #[test] + fn six_consecutive_error_epochs_terminate() { + let mut reg = MetricBandsRegistry::new( + bands_one("train_loss", band(1e-6, 10.0, 0.0, 1000.0)), + BandSettings { consecutive_epochs_for_warn: 3, consecutive_epochs_for_error: 6 }, + ); + // 5 epochs of saturated loss — streak below threshold + for i in 0..5 { + let result = reg.update_and_check("train_loss", 1e15); + assert!( + result.is_none(), + "epoch {i}: expected None (streak < 6), got {result:?}", + ); + } + // 6th consecutive epoch — must terminate + let term = reg + .update_and_check("train_loss", 1e15) + .expect("expected RegressionError on 6th out-of-band epoch"); + match term { + TerminationReason::RegressionError { metric, consecutive, .. } => { + assert_eq!(metric, "train_loss"); + assert_eq!(consecutive, 6); + } + other => panic!("expected RegressionError, got {other:?}"), + } + } + + #[test] + fn nan_treated_as_out_of_band() { + let mut reg = MetricBandsRegistry::new( + bands_one("m", band(-100.0, 100.0, -1000.0, 1000.0)), + BandSettings { consecutive_epochs_for_warn: 2, consecutive_epochs_for_error: 4 }, + ); + for _ in 0..3 { + assert!(reg.update_and_check("m", f64::NAN).is_none()); + } + let term = reg + .update_and_check("m", f64::NAN) + .expect("NaN must count as out-of-band"); + assert!(matches!(term, TerminationReason::RegressionError { .. })); + } + + #[test] + fn warn_band_does_not_terminate_alone() { + // Warn band [0, 1], error band [-100, 100]. value=5 is in warn-only. + let mut reg = MetricBandsRegistry::new( + bands_one("m", band(0.0, 1.0, -100.0, 100.0)), + BandSettings { consecutive_epochs_for_warn: 2, consecutive_epochs_for_error: 4 }, + ); + for _ in 0..20 { + assert!(reg.update_and_check("m", 5.0).is_none()); + } + } + + #[test] + fn loading_invalid_toml_rejected() { + // Construct an inverted band directly — MetricBandsRegistry::new would + // accept it, but load_from_toml's validator rejects. + let toml_text = r#" +[settings] +consecutive_epochs_for_warn = 3 +consecutive_epochs_for_error = 6 + +[bands.bad] +warn_low = 10.0 +warn_high = -10.0 +error_low = -100.0 +error_high = 100.0 +"#; + let dir = std::env::temp_dir(); + let path = dir.join("metric-bands-invalid-test.toml"); + std::fs::write(&path, toml_text).unwrap(); + let err = MetricBandsRegistry::load_from_toml(&path).unwrap_err(); + let msg = format!("{err:#}"); + assert!( + msg.contains("low > high") || msg.contains("invariant violation"), + "unexpected error message: {msg}", + ); + let _ = std::fs::remove_file(&path); + } +} diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 931ab4f18..092bede09 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -769,6 +769,68 @@ impl DQNTrainer { q_diagnostics, ).await?; + // Plan 5 Task 2 — A.4 regression-detection hard-stop. + // + // Run after the HEALTH_DIAG line was emitted by + // `log_epoch_metrics_and_financials` so the band check sees the + // same values that were just written to the log. `harvest_health_diag_metrics` + // returns a small Vec<(name, value)> of the keys defined in + // `config/metric-bands.toml` — unknown keys in the toml return + // `None` from `update_and_check` (silent OK per Invariant 7). + // + // When ANY metric crosses the `2N` consecutive-error-band threshold, + // we log a final TERMINATED_BY_REGRESSION line, emit ERROR-level + // tracing under target=`dqn::regression`, and return + // `CommonError::RegressionDetected` so the call stack unwinds + // cleanly through the trainer's `Drop` impl (which tears down CUDA + // resources). No SIGTERM raise — the error itself is the signal. + if self.metric_bands.band_count() > 0 { + let harvested = self.harvest_health_diag_metrics(epoch, &log_output); + for (name, value) in harvested { + if let Some(term) = self.metric_bands.update_and_check(&name, value) { + match term { + super::monitoring::TerminationReason::RegressionError { + metric, value, band, consecutive, + } => { + tracing::error!( + target: "dqn::regression", + epoch = epoch, + metric = %metric, + value = value, + consecutive = consecutive, + band_error_low = band.error_low, + band_error_high = band.error_high, + "TERMINATED_BY_REGRESSION: metric `{metric}` outside error band \ + [{:.6e}, {:.6e}] for {consecutive} consecutive epochs (value={:.6e})", + band.error_low, band.error_high, value, + ); + tracing::info!( + "HEALTH_DIAG[{}]: TERMINATED_BY_REGRESSION metric={} value={:.6e} \ + consecutive={} band_error=[{:.6e}, {:.6e}]", + epoch, metric, value, consecutive, band.error_low, band.error_high, + ); + return Err(anyhow::Error::new( + common::error::CommonError::RegressionDetected { + metric, + value, + band_warn_low: band.warn_low, + band_warn_high: band.warn_high, + band_error_low: band.error_low, + band_error_high: band.error_high, + consecutive, + }, + )); + } + super::monitoring::TerminationReason::RegressionWarn { .. } => { + // Warn variant currently never returned by + // update_and_check (warns only log). Kept for + // exhaustive matching. + } + } + } + } + } + total_loss += log_output.avg_loss; total_q_value += log_output.avg_q_value; total_gradient_norm += log_output.avg_grad_norm; @@ -4726,5 +4788,100 @@ impl DQNTrainer { ); Ok(()) } + + /// Plan 5 Task 2 — collect the per-epoch HEALTH_DIAG metrics that the + /// regression-detection registry should evaluate. + /// + /// Returns a `Vec<(metric_name, value)>` whose names MUST match the keys + /// in `config/metric-bands.toml` exactly. The set is deliberately a + /// subset of the full HEALTH_DIAG output — only the values that are + /// stable across the warmup → steady-state transition and that we have + /// known-good baseline distributions for. Adding a new key here without + /// also adding a band to the toml is harmless (the registry returns + /// `None` for unknown keys). + /// + /// All values are read from `self` / `log_output` directly — no GPU + /// readback is performed (the values were already harvested by + /// `log_epoch_metrics_and_financials` and the HEALTH_DIAG block above). + pub(crate) fn harvest_health_diag_metrics( + &self, + _epoch: usize, + log_output: &EpochLogOutput, + ) -> Vec<(String, f64)> { + let mut out: Vec<(String, f64)> = Vec::with_capacity(16); + + // ── Core training metrics (load-bearing for stability) ── + out.push(("train_loss".to_owned(), log_output.avg_loss)); + out.push(("avg_q_value".to_owned(), log_output.avg_q_value)); + out.push(("avg_grad_norm".to_owned(), log_output.avg_grad_norm)); + out.push(("train_sharpe".to_owned(), log_output.epoch_sharpe)); + // val_loss is `-sharpe` per the existing convention; surface as + // `val_sharpe` for the band file (more intuitive band semantics). + out.push(("val_sharpe".to_owned(), -log_output.val_loss)); + out.push(("q_min".to_owned(), log_output.q_min)); + out.push(("q_max".to_owned(), log_output.q_max)); + out.push(("q_mean".to_owned(), log_output.q_mean)); + + // ── Effective controller outputs (cached during HEALTH_DIAG block) ── + if let Some(v) = self.last_cql_alpha_eff { + out.push(("cql_alpha_eff".to_owned(), v as f64)); + } + if let Some(v) = self.last_iqn_budget_eff { + out.push(("iqn_budget_eff".to_owned(), v as f64)); + } + if let Some(v) = self.last_c51_budget_eff { + out.push(("c51_budget_eff".to_owned(), v as f64)); + } + if let Some(v) = self.last_tau_eff { + out.push(("tau_eff".to_owned(), v as f64)); + } + if let Some(v) = self.last_gamma_eff { + out.push(("gamma_eff".to_owned(), v as f64)); + } + + // ── Diagnostic signals from HEALTH_DIAG diag/explore blocks ── + out.push(("training_sharpe_ema".to_owned(), self.training_sharpe_ema as f64)); + if let Some(v) = self.last_action_entropy { + out.push(("action_entropy".to_owned(), v as f64)); + } + + // ── ISV slot EMAs that have controllers we want to band-monitor ── + // Read once per epoch via the cold-path ISV accessor; matches the + // values logged in HEALTH_DIAG. + if let Some(ref fused) = self.fused_ctx { + use crate::cuda_pipeline::gpu_dqn_trainer::{ + AUX_NEXT_BAR_MSE_EMA_INDEX, AUX_REGIME_CE_EMA_INDEX, + REWARD_POPART_EMA_INDEX, REWARD_CF_EMA_INDEX, + REWARD_TRAIL_EMA_INDEX, REWARD_MICRO_EMA_INDEX, + }; + let trainer = fused.trainer(); + out.push(( + "aux_next_bar_mse_ema".to_owned(), + trainer.read_isv_signal_at(AUX_NEXT_BAR_MSE_EMA_INDEX) as f64, + )); + out.push(( + "aux_regime_ce_ema".to_owned(), + trainer.read_isv_signal_at(AUX_REGIME_CE_EMA_INDEX) as f64, + )); + out.push(( + "isv_reward_popart_ema".to_owned(), + trainer.read_isv_signal_at(REWARD_POPART_EMA_INDEX) as f64, + )); + out.push(( + "isv_reward_cf_ema".to_owned(), + trainer.read_isv_signal_at(REWARD_CF_EMA_INDEX) as f64, + )); + out.push(( + "isv_reward_trail_ema".to_owned(), + trainer.read_isv_signal_at(REWARD_TRAIL_EMA_INDEX) as f64, + )); + out.push(( + "isv_reward_micro_ema".to_owned(), + trainer.read_isv_signal_at(REWARD_MICRO_EMA_INDEX) as f64, + )); + } + + out + } } diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index 8392dfa09..762ba0b93 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -43,6 +43,10 @@ | `trainers/dqn/trainer/metrics.rs` | `training_loop.rs` | Wired | Epoch-boundary metric compilation + backtest eval | — | | `trainers/dqn/trainer/state.rs` | `trainer/mod.rs` | Wired | Mutable training state | — | | `trainers/dqn/trainer/training_loop.rs` | called from `DQNTrainer::train()` | Wired | Main epoch loop | — | +| `trainers/dqn/trainer/monitoring.rs` (`MetricBandsRegistry`, `MetricBands`, `BandSettings`, `TerminationReason`) | `trainer/training_loop.rs` (post-HEALTH_DIAG harvest+check), `trainer/constructor.rs` (auto-load `config/metric-bands.toml`), `smoke_tests/regression_detection.rs` | Wired | Plan 5 Task 2 A.4 — regression-detection hard-stop (2N consecutive error-band epochs → `CommonError::RegressionDetected`); cold-path (per-epoch); zero hot-path overhead | — | +| `crates/common/src/error.rs` (`CommonError::RegressionDetected`) | returned from `training_loop::train_with_data_full_loop_slices`; consumed by `services/trading_service/src/error.rs::From` | Wired | Plan 5 Task 2 — terminal training error variant; flat band fields (no upward dep on ml crate) | — | +| `config/metric-bands.toml` | `MetricBandsRegistry::load_from_toml` (auto-loaded by `DQNTrainer::new_internal`) | Wired | Plan 5 Task 2 — per-metric warn/error bands (mean ± 5σ / mean ± 10σ); populated by `scripts/populate-metric-bands-from-runs.py` from Plan 5 Task 1B aggregate JSONs | — | +| `scripts/populate-metric-bands-from-runs.py` | reads aggregate JSONs from `scripts/aggregate-multi-seed-metrics.py` (Plan 5 Task 1B.2); writes `config/metric-bands.toml` `[bands]` section | Wired | Plan 5 Task 2 helper — N>=3 sample path; N=1 multiplicative fallback | — | | `trainers/dqn/adaptive_monitor.rs` | Read-only observer trait + harness (FireRateStats, DiagSnapshot, IsvBus<'a>); consumers added in Plan 1 Tasks 9-17 (atoms/gamma/kelly_cap/tau/epsilon/grad_balancer monitors) | Wired (consumers added in same plan) | C.6 GPU-drives-CPU-reads | — | | `trainers/dqn/monitors/grad_balancer_monitor.rs` | Read-only observer for grad_balance_isv_update kernel output (ISV slots 31..35); consumers: HEALTH_DIAG + controller_activity smoke | Wired | Plan 1 Task 17 | — | | `trainers/dqn/monitors/tau_monitor.rs` | Read-only observer for tau_update kernel output (ISV slot 42); consumers: HEALTH_DIAG + controller_activity smoke | Wired | Plan 1 Task 13 | — | diff --git a/scripts/populate-metric-bands-from-runs.py b/scripts/populate-metric-bands-from-runs.py new file mode 100755 index 000000000..55447dab4 --- /dev/null +++ b/scripts/populate-metric-bands-from-runs.py @@ -0,0 +1,328 @@ +#!/usr/bin/env python3 +"""Plan 5 Task 2 — populate `config/metric-bands.toml` from N known-good aggregate JSONs. + +Consumes the per-stream aggregate JSONs produced by +`scripts/aggregate-multi-seed-metrics.py` (Plan 5 Task 1B.2) and emits TOML +band entries at `mean ± 5σ` (warn) / `mean ± 10σ` (error). The N>=2 case is the +production path; for N=1 we fall back to safe symmetric multiplicative +defaults documented inline. + +Usage: + + populate-metric-bands-from-runs.py \\ + --aggregate run-a.json --aggregate run-b.json --aggregate run-c.json \\ + --warmup-end-epoch 10 \\ + --include-metrics avg_q_value train_loss val_sharpe ... \\ + --output config/metric-bands.toml + +Each `--aggregate` JSON is expected to contain the schema written by +`aggregate-multi-seed-metrics.py`: + + { "aggregates": { "": [ {"epoch": int, "mean": float, ...}, ... ], ... } } + +The script computes per-metric POST-WARMUP statistics (epoch >= warmup-end) +across the union of all samples from all aggregate JSONs: + + band_warn = [grand_mean - 5σ, grand_mean + 5σ] + band_error = [grand_mean - 10σ, grand_mean + 10σ] + +Where σ is the population standard deviation of the per-(epoch, run) sample +values (NOT the std of per-run means — we want to capture WITHIN-RUN epoch +drift in the band, not just CROSS-RUN dispersion). + +For metrics where the variance is genuinely zero post-warmup (e.g. +`tau_eff = 0.005` always), we widen the band slightly so a single out-of- +band epoch doesn't trip on numerical noise: + + σ_eff = max(σ, |grand_mean| * 1e-4 + 1e-9) + +For N=1 sample we fall back to multiplicative defaults — same shape as the +plan's `[settings]` block but multiplicative since we have no sigma: + + band_warn = [|x| * 0.5, |x| * 2.0] (signed: [min(x*0.5,x*2), max(x*0.5,x*2)]) + band_error = [|x| * 0.1, |x| * 10.0] + +Existing bands in the output toml are PRESERVED unless --force is passed. +The `[settings]` block is preserved verbatim. +""" +from __future__ import annotations + +import argparse +import json +import math +import sys +from collections import defaultdict +from pathlib import Path +from statistics import pstdev + + +def load_aggregate_json(path: Path) -> dict: + """Load an aggregate JSON written by aggregate-multi-seed-metrics.py.""" + with path.open("r", encoding="utf-8") as fh: + return json.load(fh) + + +def collect_post_warmup_samples( + aggregates: list[dict], + warmup_end_epoch: int, +) -> dict[str, list[float]]: + """For each metric, collect per-epoch samples from epoch >= warmup_end across all runs. + + Each aggregate's `aggregates.[i].mean` already represents the + cross-seed mean at epoch i. We keep that as a single sample (one per + run × epoch). Cross-run dispersion is captured by including all runs' + means; within-run epoch drift is captured by including all post-warmup + epochs. + """ + by_metric: dict[str, list[float]] = defaultdict(list) + for agg in aggregates: + for metric_name, rows in agg.get("aggregates", {}).items(): + for row in rows: + if row.get("epoch", -1) >= warmup_end_epoch: + val = row.get("mean") + if val is not None and math.isfinite(val): + by_metric[metric_name].append(float(val)) + return dict(by_metric) + + +def compute_bands( + samples: list[float], + warn_sigma: float = 5.0, + error_sigma: float = 10.0, +) -> dict[str, float] | None: + """Compute warn / error bands from N samples. Returns None if N==0.""" + n = len(samples) + if n == 0: + return None + mean_val = sum(samples) / n + if n == 1: + # Multiplicative fallback — captures sign while widening generously. + x = samples[0] + # Symmetric multiplicative shape; collapse to additive for x near 0. + if abs(x) < 1e-9: + return { + "warn_low": -1.0, "warn_high": 1.0, + "error_low": -10.0, "error_high": 10.0, + } + warn_lo = min(x * 0.5, x * 2.0) + warn_hi = max(x * 0.5, x * 2.0) + err_lo = min(x * 0.1, x * 10.0) + err_hi = max(x * 0.1, x * 10.0) + # Make sure error envelope strictly contains warn envelope. + err_lo = min(err_lo, warn_lo - abs(warn_hi - warn_lo)) + err_hi = max(err_hi, warn_hi + abs(warn_hi - warn_lo)) + return { + "warn_low": warn_lo, "warn_high": warn_hi, + "error_low": err_lo, "error_high": err_hi, + } + sigma = pstdev(samples) + # Effective sigma — never zero. + sigma_eff = max(sigma, abs(mean_val) * 1e-4 + 1e-9) + return { + "warn_low": mean_val - warn_sigma * sigma_eff, + "warn_high": mean_val + warn_sigma * sigma_eff, + "error_low": mean_val - error_sigma * sigma_eff, + "error_high": mean_val + error_sigma * sigma_eff, + } + + +def format_toml( + settings_block: str, + bands: dict[str, dict[str, float]], +) -> str: + """Render a metric-bands.toml document with the given settings + bands. + + Settings block is taken verbatim from the existing file (or a sensible + default). Bands are written under a flat `[bands]` table with one + sub-table per metric: + + [bands.] + warn_low = ... + warn_high = ... + error_low = ... + error_high = ... + """ + lines: list[str] = [] + lines.append("# Regression-detection bands for DQN training metrics.") + lines.append("# Invariant 7. Per A.4, bands populated from last 3 known-good runs.") + lines.append("# Band = [mean - 5σ, mean + 5σ] per metric (warn);") + lines.append("# [mean - 10σ, mean + 10σ] (error). Recomputed after baseline promotion.") + lines.append("# N consecutive out-of-band epochs → WARN; 2N consecutive → ERROR + clean termination.") + lines.append("") + lines.append(settings_block.rstrip()) + lines.append("") + lines.append("[bands]") + lines.append("") + for metric in sorted(bands.keys()): + b = bands[metric] + lines.append(f"[bands.{metric}]") + lines.append(f"warn_low = {_fmt_float(b['warn_low'])}") + lines.append(f"warn_high = {_fmt_float(b['warn_high'])}") + lines.append(f"error_low = {_fmt_float(b['error_low'])}") + lines.append(f"error_high = {_fmt_float(b['error_high'])}") + lines.append("") + return "\n".join(lines).rstrip() + "\n" + + +def _fmt_float(x: float) -> str: + """TOML-friendly float rendering with reasonable precision.""" + if not math.isfinite(x): + # toml has no NaN/Inf; use a sentinel that downstream load-time + # validation will reject loudly. + return "0.0 # WARNING: non-finite computed band — please check inputs" + if x == 0.0: + return "0.0" + if abs(x) >= 1e6 or abs(x) < 1e-3: + return f"{x:.6e}" + return f"{x:.6f}" + + +def parse_existing_settings(path: Path) -> str: + """Extract the `[settings]` block from the existing toml (verbatim).""" + if not path.is_file(): + return ( + "[settings]\n" + "consecutive_epochs_for_warn = 3\n" + "consecutive_epochs_for_error = 6\n" + ) + text = path.read_text(encoding="utf-8") + out_lines: list[str] = [] + in_block = False + for line in text.splitlines(): + stripped = line.strip() + if stripped == "[settings]": + in_block = True + out_lines.append(line) + continue + if in_block: + if stripped.startswith("[") and stripped.endswith("]"): + break + out_lines.append(line) + if not out_lines: + return ( + "[settings]\n" + "consecutive_epochs_for_warn = 3\n" + "consecutive_epochs_for_error = 6\n" + ) + return "\n".join(out_lines) + "\n" + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + ap.add_argument( + "--aggregate", + action="append", + type=Path, + default=[], + help="Path to an aggregate JSON (output of aggregate-multi-seed-metrics.py). " + "Pass once per known-good run. N>=3 is the plan-prescribed minimum; N=1 falls back to safe defaults.", + ) + ap.add_argument( + "--warmup-end-epoch", + type=int, + default=10, + help="Epoch at which warmup ends — only epoch >= this contributes to band statistics. Default 10.", + ) + ap.add_argument( + "--include-metrics", + nargs="+", + default=[], + help="If non-empty, restrict band emission to these metric names only. Otherwise emit all metrics found.", + ) + ap.add_argument( + "--output", + type=Path, + required=True, + help="Output TOML path (typically config/metric-bands.toml).", + ) + ap.add_argument( + "--force", + action="store_true", + help="Overwrite existing band entries instead of preserving them. Settings block is always preserved.", + ) + args = ap.parse_args() + + if not args.aggregate: + print("ERROR: at least one --aggregate JSON is required", file=sys.stderr) + return 1 + + aggregates = [] + for p in args.aggregate: + if not p.is_file(): + print(f"ERROR: aggregate JSON not found: {p}", file=sys.stderr) + return 1 + aggregates.append(load_aggregate_json(p)) + + by_metric = collect_post_warmup_samples(aggregates, args.warmup_end_epoch) + + if args.include_metrics: + keep = set(args.include_metrics) + by_metric = {k: v for k, v in by_metric.items() if k in keep} + + bands: dict[str, dict[str, float]] = {} + for metric_name, samples in by_metric.items(): + b = compute_bands(samples) + if b is not None: + bands[metric_name] = b + + if not bands: + print( + "ERROR: no bands computable — check warmup-end-epoch and metric names " + "(post-warmup samples may be empty)", + file=sys.stderr, + ) + return 1 + + settings_block = parse_existing_settings(args.output) + + # Preserve existing bands unless --force. + if not args.force and args.output.is_file(): + existing_text = args.output.read_text(encoding="utf-8") + # Naive parse — we look for `[bands.]` headers and keep the + # block intact if we're not generating a new one for that metric. + # Full parse is overkill since we control the writer too. + existing_blocks = _parse_band_blocks(existing_text) + for metric, body in existing_blocks.items(): + if metric not in bands: + bands[metric] = body + + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(format_toml(settings_block, bands), encoding="utf-8") + + print( + f"Wrote {len(bands)} band entries to {args.output} " + f"(from {len(args.aggregate)} aggregate JSON(s), " + f"warmup-end={args.warmup_end_epoch})", + file=sys.stderr, + ) + return 0 + + +def _parse_band_blocks(text: str) -> dict[str, dict[str, float]]: + """Naive parser for `[bands.]` sub-tables. Returns dict by metric name.""" + out: dict[str, dict[str, float]] = {} + current_metric: str | None = None + current_body: dict[str, float] = {} + for line in text.splitlines(): + stripped = line.strip() + if stripped.startswith("[bands.") and stripped.endswith("]"): + if current_metric is not None and current_body: + out[current_metric] = current_body + current_metric = stripped[len("[bands."):-1] + current_body = {} + continue + if current_metric is not None and "=" in stripped and not stripped.startswith("#"): + key, _, value = stripped.partition("=") + key = key.strip() + value = value.split("#", 1)[0].strip() + try: + current_body[key] = float(value) + except ValueError: + pass + if current_metric is not None and current_body: + out[current_metric] = current_body + return out + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/services/trading_service/src/error.rs b/services/trading_service/src/error.rs index 0000c2c52..651c00ee2 100644 --- a/services/trading_service/src/error.rs +++ b/services/trading_service/src/error.rs @@ -105,6 +105,16 @@ impl From for tonic::Status { actual_ms, max_ms )) }, + common::error::CommonError::RegressionDetected { + ref metric, + value, + consecutive, + .. + } => tonic::Status::failed_precondition(format!( + "ML training regression detected: metric={} value={:.6e} \ + consecutive_epochs={}", + metric, value, consecutive + )), } }, TradingServiceError::OrderValidation { reason } => {