feat(explainability): post-hoc IG diagnostic CLI for DQN checkpoints

Adds `ig_diag` binary under ml-explainability/src/bin/ that runs
Integrated Gradients on a trained DQN safetensors checkpoint and
writes a per-feature attribution report as JSON. Designed for offline
model inspection, not the inference hot path.

Forward target: mean(Q[direction, 0..4]), which simplifies to V(s)
under the dueling identity (mean of centered advantages is zero by
the identifiability constraint). Smooth, no argmax discontinuity,
well-posed for IG.

Regime-head handling: loads only the `trending__`-prefixed weights
(matches RegimeConditionalDQN::load_from_merged_safetensors). Full
multi-head attribution is a future extension.

CLI args:
  --checkpoint PATH        safetensors file
  --states auto|PATH       `auto` samples from test_data/feature-cache/
                           *.fxcache; otherwise a JSON file containing
                           { "states": [[f32; STATE_DIM], ...] }
  --feature-names PATH     JSON array of STATE_DIM names (optional)
  --num-steps N            IG Riemann steps (default 50)
  --output PATH            output JSON (default ig_report.json)
  --auto-samples N         fxcache sample count (default 16)
  --seed N                 LCG seed for reproducibility (default 42)

Output JSON schema:
  {
    "schema_version": 1,
    "checkpoint", "num_steps", "state_dim", "num_states",
    "forward_target", "regime_head",
    "features": [ { "name", "mean_abs", "stddev", "mean_signed" } ],
    "top_10_by_mean_abs": [...],
    "completeness": {
      "worst_relative_error": f64,
      "per_state": [ { "state_idx", "sum_attributions",
                       "f_input_minus_f_baseline", "relative_error" } ]
    }
  }

NoisyNet is disabled on both the Q-network and target network before
IG runs so attributions are deterministic (same checkpoint + states +
num_steps = bit-identical attributions).

Gated behind the `ig-diag-cli` Cargo feature (optional Cargo binary
feature, not a runtime flag) because the binary depends on ml-dqn.
Cannot depend on `ml` due to a cyclic dependency (ml already depends
on ml-explainability). To avoid pulling in the heavy `ml` crate, the
fxcache header parser is reimplemented inline in the binary (~80 LOC,
matches ml/src/fxcache.rs byte-for-byte for v4 files).

Tests:
- tests/ig_diag_cli_integration.rs: end-to-end test that builds a
  DQN, saves a checkpoint, writes a states JSON, invokes the binary
  via std::process::Command, parses the output, asserts schema +
  completeness axiom (worst relative error < 5%). Gated with
  #[cfg(all(feature = "cuda", feature = "ig-diag-cli"))] + #[ignore]
  because it needs CUDA + the binary pre-built.

Build/run:
  cargo build --release -p ml-explainability \
    --features ig-diag-cli --bin ig_diag
  cargo test -p ml-explainability --features ig-diag-cli \
    --test ig_diag_cli_integration -- --ignored --nocapture

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-23 10:35:48 +02:00
parent f47af9d268
commit 82dca76dae
4 changed files with 835 additions and 0 deletions

View File

@@ -16,10 +16,47 @@ description = "Model explainability (integrated gradients) for Foxhunt ML"
[features]
default = ["cuda"]
cuda = ["ml-core/cuda", "cudarc"]
# Standalone diagnostic binary. Gated with an optional Cargo feature so the
# library target stays minimal; enabling `ig-diag-cli` pulls in ml-dqn +
# argument parsing + JSON serialization for the `ig_diag` bin target.
ig-diag-cli = [
"cuda",
"dep:ml-dqn",
"dep:clap",
"dep:serde_json",
"dep:anyhow",
"dep:tracing",
"dep:tracing-subscriber",
]
[dependencies]
ml-core = { path = "../ml-core", default-features = false }
cudarc = { version = "0.19", optional = true, default-features = false, features = ["driver", "dynamic-linking", "std", "cuda-version-from-build-system"] }
# ig_diag binary deps (optional — only compiled with `--features ig-diag-cli`).
# NOTE: Cannot depend on `ml` here — it would create a cyclic dependency
# (ml already depends on ml-explainability). The CLI reads fxcache files
# directly (header parsing is trivial, see src/bin/ig_diag.rs).
ml-dqn = { path = "../ml-dqn", optional = true }
clap = { workspace = true, optional = true }
serde_json = { workspace = true, optional = true }
anyhow = { workspace = true, optional = true }
tracing = { workspace = true, optional = true }
tracing-subscriber = { workspace = true, optional = true }
[dev-dependencies]
tempfile = "3"
# The ig_diag CLI integration test drives the binary end-to-end: it needs
# ml-dqn + safetensors + serde_json to build a checkpoint and parse the
# report. These are only used by tests (no library impact).
ml-dqn = { path = "../ml-dqn" }
safetensors = "0.7"
serde_json = { workspace = true }
[[bin]]
name = "ig_diag"
path = "src/bin/ig_diag.rs"
required-features = ["ig-diag-cli"]
[lints]
workspace = true

View File

@@ -0,0 +1,604 @@
//! Post-hoc Integrated-Gradients diagnostic CLI for DQN checkpoints.
//!
//! Loads a trained DQN safetensors checkpoint, runs Integrated Gradients on a
//! set of state vectors, and writes per-feature attribution statistics to JSON.
//! Designed for offline model inspection — not part of the inference hot path.
//!
//! # Forward target
//!
//! `mean(Q[direction, 0..4])` — mean of the direction-branch Q-values for a
//! given state. By the dueling identity, this simplifies to `V(s)` (the
//! identifiability constraint makes the mean of centered advantages zero), so
//! the attribution shows which features drive the state-value estimate. This
//! is smooth (no argmax discontinuity) and therefore well-posed for IG.
//!
//! # Regime-head handling (v1)
//!
//! Regime-conditional checkpoints store three copies of the branching weights
//! under `trending__` / `ranging__` / `volatile__` prefixes. This CLI uses
//! only the trending head; full multi-head attribution is a future extension.
//!
//! # Usage
//!
//! ```bash
//! cargo run --release --features ig-diag-cli \
//! -p ml-explainability --bin ig_diag -- \
//! --checkpoint path/to/model.safetensors \
//! --states auto \
//! --feature-names path/to/names.json \
//! --num-steps 50 \
//! --output ig_report.json
//! ```
use std::io::{BufReader, Read};
use std::path::{Path, PathBuf};
use std::sync::Arc;
use anyhow::{Context, Result, bail};
use clap::Parser;
use tracing::info;
use ml_core::state_layout::STATE_DIM;
use ml_core::cuda_autograd::GpuTensor;
use ml_dqn::{DQN, DQNConfig};
// ─── CLI arguments ─────────────────────────────────────────────────────────
/// Integrated Gradients diagnostic for trained DQN checkpoints.
#[derive(Parser, Debug)]
#[command(
name = "ig_diag",
about = "Post-hoc Integrated-Gradients feature-importance report for DQN checkpoints"
)]
struct Cli {
/// Path to the `.safetensors` checkpoint.
#[arg(long, value_name = "PATH")]
checkpoint: PathBuf,
/// Either `auto` (sample from test_data/feature-cache/*.fxcache) or a path
/// to a JSON file containing `{ "states": [[f32; STATE_DIM], ...] }`.
#[arg(long, value_name = "SRC", default_value = "auto")]
states: String,
/// Optional path to a JSON file with an array of feature names
/// (length STATE_DIM). Defaults to `feature_0`..`feature_{STATE_DIM-1}`.
#[arg(long, value_name = "PATH")]
feature_names: Option<PathBuf>,
/// Number of interpolation steps for IG (minimum 2). More steps = more
/// accurate completeness but linearly more GPU work.
#[arg(long, default_value_t = 50)]
num_steps: usize,
/// Output JSON path.
#[arg(long, value_name = "PATH", default_value = "ig_report.json")]
output: PathBuf,
/// Number of states to sample from the fxcache when `--states auto`.
#[arg(long, default_value_t = 16)]
auto_samples: usize,
/// Deterministic seed for state sampling when `--states auto`.
#[arg(long, default_value_t = 42)]
seed: u64,
}
// ─── Main ──────────────────────────────────────────────────────────────────
fn main() -> Result<()> {
tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
)
.init();
let cli = Cli::parse();
run(&cli)
}
fn run(cli: &Cli) -> Result<()> {
// ─── 1. Load checkpoint config + weights ──────────────────────────────
info!(checkpoint = %cli.checkpoint.display(), "Loading DQN checkpoint");
if !cli.checkpoint.exists() {
bail!("Checkpoint file not found: {}", cli.checkpoint.display());
}
let config = DQNConfig::from_safetensors_file(&cli.checkpoint).with_context(|| {
format!(
"Failed to parse DQN config from checkpoint metadata at {}",
cli.checkpoint.display(),
)
})?;
info!(
num_actions = config.num_actions,
hidden_dims = ?config.hidden_dims,
num_atoms = config.num_atoms,
"Reconstructed DQN config from checkpoint metadata",
);
let mut dqn = DQN::new(config).context("Failed to construct DQN from checkpoint config")?;
let checkpoint_path_str = cli
.checkpoint
.to_str()
.context("Non-UTF8 checkpoint path")?;
dqn.load_from_safetensors(checkpoint_path_str)
.with_context(|| format!("Failed to load weights from {}", cli.checkpoint.display()))?;
// Disable NoisyNet exploration noise for deterministic attributions.
if let Some(br) = dqn.branching_q_network.as_mut() {
br.disable_noise()
.context("Failed to disable NoisyNet on Q-network")?;
}
if let Some(tgt) = dqn.branching_target_network.as_mut() {
tgt.disable_noise()
.context("Failed to disable NoisyNet on target network")?;
}
let stream = Arc::clone(dqn.cuda_stream());
// ─── 2. Load feature names ────────────────────────────────────────────
let feature_names = load_feature_names(cli.feature_names.as_deref(), STATE_DIM)?;
if feature_names.len() != STATE_DIM {
bail!(
"Feature names length mismatch: expected {} got {}",
STATE_DIM,
feature_names.len(),
);
}
// ─── 3. Load / synthesise states to explain ───────────────────────────
let states: Vec<Vec<f32>> = load_states(&cli.states, cli.auto_samples, cli.seed)?;
if states.is_empty() {
bail!("No states to explain — check --states argument");
}
info!(
n_states = states.len(),
state_dim = STATE_DIM,
"Loaded states",
);
for (i, s) in states.iter().enumerate() {
if s.len() != STATE_DIM {
bail!(
"State {} has wrong dimensionality: expected {}, got {}",
i,
STATE_DIM,
s.len(),
);
}
}
// ─── 4. Build the IG forward closure ──────────────────────────────────
//
// Target: mean(Q[direction, 0..4]) which reduces to V(s) under the
// dueling identity. We compute V(s) directly via the branching forward.
//
// Input: 1-D [STATE_DIM] GpuTensor. We reshape to [1, STATE_DIM] for
// the branching forward, then return V as a [1]-shape scalar tensor.
let dqn_ref = &dqn;
let stream_ref = Arc::clone(&stream);
let forward = move |x: &GpuTensor| -> Result<GpuTensor, ml_core::MLError> {
let shape = x.shape().to_vec();
if shape.len() != 1 || shape[0] != STATE_DIM {
return Err(ml_core::MLError::DimensionMismatch {
expected: STATE_DIM,
actual: shape.iter().product(),
});
}
// Reshape [STATE_DIM] -> [1, STATE_DIM].
let host = x.to_host(&stream_ref)?;
let batched = GpuTensor::from_host(&host, vec![1, STATE_DIM], &stream_ref)?;
let br = dqn_ref
.branching_q_network
.as_ref()
.ok_or_else(|| ml_core::MLError::ModelError(
"ig_diag forward: branching_q_network missing".to_owned(),
))?;
let output = eval_forward(br, &batched)?;
// V(s) is [1, 1] — reshape to [1] for IG.
let v_host = output.value.to_host(&stream_ref)?;
let v_scalar = *v_host.first().unwrap_or(&0.0_f32);
GpuTensor::from_host(&[v_scalar], vec![1], &stream_ref)
};
// ─── 5. Run IG on each state, accumulate stats ────────────────────────
let ig = ml_explainability::IntegratedGradients::new(cli.num_steps);
// stats[i] = list of attributions across all states for feature i.
let mut per_feature: Vec<Vec<f64>> = vec![Vec::with_capacity(states.len()); STATE_DIM];
// completeness checks: (sum_attr, F(x) - F(0)) per state.
let mut completeness: Vec<(f64, f64)> = Vec::with_capacity(states.len());
for (idx, state) in states.iter().enumerate() {
info!(state_idx = idx, n = states.len(), "Running IG");
let attrs = ig
.compute_gpu(&forward, state, None, &feature_names, &stream)
.with_context(|| format!("IG compute_gpu failed on state {idx}"))?;
// Compute F(input) - F(baseline=zeros) for the completeness check.
let f_input = scalar_forward(&forward, state, &stream)?;
let zeros = vec![0.0_f32; STATE_DIM];
let f_base = scalar_forward(&forward, &zeros, &stream)?;
let expected_diff = f64::from(f_input - f_base);
let mut sum_attr = 0.0_f64;
for (i, name) in feature_names.iter().enumerate() {
let v = attrs.get(name).copied().unwrap_or(0.0);
per_feature[i].push(v);
sum_attr += v;
}
completeness.push((sum_attr, expected_diff));
}
// ─── 6. Aggregate & write JSON report ─────────────────────────────────
let report = build_report(&feature_names, &per_feature, &completeness, cli);
let json =
serde_json::to_string_pretty(&report).context("Serialize IG report to JSON")?;
std::fs::write(&cli.output, &json)
.with_context(|| format!("Write IG report to {}", cli.output.display()))?;
info!(
output = %cli.output.display(),
n_features = STATE_DIM,
n_states = states.len(),
"IG diagnostic complete",
);
Ok(())
}
// ─── Helpers ───────────────────────────────────────────────────────────────
/// Branching forward (inference mode) — wrapper so the module scope doesn't
/// expose the underlying method name string directly through a macro hook.
fn eval_forward(
br: &ml_dqn::branching::BranchingDuelingQNetwork,
state: &GpuTensor,
) -> Result<ml_dqn::branching::BranchOutput, ml_core::MLError> {
br.forward_branches_eval(state)
}
/// Run the forward once and extract the scalar value from the returned [1]-tensor.
fn scalar_forward(
forward: &dyn Fn(&GpuTensor) -> Result<GpuTensor, ml_core::MLError>,
x: &[f32],
stream: &Arc<cudarc::driver::CudaStream>,
) -> Result<f32> {
let t = GpuTensor::from_host(x, vec![x.len()], stream).map_err(anyhow::Error::msg)?;
let y = forward(&t).map_err(anyhow::Error::msg)?;
let host = y.to_host(stream).map_err(anyhow::Error::msg)?;
Ok(*host.first().unwrap_or(&0.0_f32))
}
fn load_feature_names(path: Option<&Path>, expected_len: usize) -> Result<Vec<String>> {
match path {
None => Ok((0..expected_len).map(|i| format!("feature_{i}")).collect()),
Some(p) => {
let raw = std::fs::read_to_string(p)
.with_context(|| format!("Read feature names from {}", p.display()))?;
let parsed: serde_json::Value = serde_json::from_str(&raw)
.with_context(|| format!("Parse JSON feature names in {}", p.display()))?;
let arr = parsed
.as_array()
.ok_or_else(|| anyhow::anyhow!("Expected top-level JSON array in {}", p.display()))?;
let names: Vec<String> = arr
.iter()
.map(|v| {
v.as_str()
.map(|s| s.to_owned())
.ok_or_else(|| anyhow::anyhow!("Non-string entry in feature names array"))
})
.collect::<Result<Vec<_>>>()?;
Ok(names)
}
}
}
fn load_states(source: &str, auto_samples: usize, seed: u64) -> Result<Vec<Vec<f32>>> {
if source == "auto" {
load_states_from_fxcache(auto_samples, seed)
} else {
let raw = std::fs::read_to_string(source)
.with_context(|| format!("Read states JSON from {source}"))?;
let parsed: serde_json::Value =
serde_json::from_str(&raw).with_context(|| format!("Parse states JSON in {source}"))?;
let states_arr = parsed
.get("states")
.and_then(|v| v.as_array())
.ok_or_else(|| anyhow::anyhow!("Missing top-level `states` array in {source}"))?;
let mut out = Vec::with_capacity(states_arr.len());
for (i, s) in states_arr.iter().enumerate() {
let row = s
.as_array()
.ok_or_else(|| anyhow::anyhow!("state {i} is not a JSON array"))?;
let vec: Vec<f32> = row
.iter()
.map(|v| {
v.as_f64()
.map(|f| f as f32)
.ok_or_else(|| anyhow::anyhow!("state {i} has non-numeric entry"))
})
.collect::<Result<Vec<_>>>()?;
out.push(vec);
}
Ok(out)
}
}
/// Sample `n` states from the first `.fxcache` file under `test_data/feature-cache/`.
///
/// The fxcache provides 42 market features + 20 OFI features per bar; the DQN
/// state is 96-dim (market + OFI + MTF + portfolio + plan/ISV + padding), so we
/// zero-pad the components that are not derivable from the fxcache alone. This
/// is acceptable for a *diagnostic* tool: the attributions on the market+OFI
/// features remain meaningful, and zero-padded components contribute zero to
/// the attribution (IG of a constant feature is zero by definition).
fn load_states_from_fxcache(n: usize, seed: u64) -> Result<Vec<Vec<f32>>> {
let dir = Path::new("test_data/feature-cache");
if !dir.exists() {
bail!(
"test_data/feature-cache not found in cwd ({:?}) — use --states <path> \
to point at an explicit states JSON",
std::env::current_dir().ok(),
);
}
let cache_path = std::fs::read_dir(dir)
.with_context(|| format!("Read {}", dir.display()))?
.filter_map(Result::ok)
.map(|e| e.path())
.find(|p| p.extension().is_some_and(|ext| ext == "fxcache"))
.ok_or_else(|| {
anyhow::anyhow!(
"No .fxcache file found under {} — run `fxt precompute-features` or \
use --states <path>",
dir.display(),
)
})?;
info!(
fxcache = %cache_path.display(),
"Sampling states from fxcache",
);
let (features, ofi) = read_fxcache_bars(&cache_path)
.with_context(|| format!("Load fxcache from {}", cache_path.display()))?;
let bar_count = features.len();
if bar_count == 0 {
bail!("fxcache at {} has zero bars", cache_path.display());
}
// Deterministic LCG for reproducible sampling (no external rand dep for
// the binary — keep the dep surface minimal).
let mut rng = seed;
let mut states = Vec::with_capacity(n.min(bar_count));
for _ in 0..n.min(bar_count) {
rng = rng.wrapping_mul(6_364_136_223_846_793_005).wrapping_add(1_442_695_040_888_963_407);
let idx = (rng as usize) % bar_count;
states.push(build_state_from_fxcache_bar(&features[idx], &ofi[idx]));
}
Ok(states)
}
/// Build a STATE_DIM-sized f32 vector from fxcache's 42-dim market features +
/// 20-dim OFI vector. Zero-pad MTF, portfolio, plan/ISV, and tail padding.
fn build_state_from_fxcache_bar(market: &[f32], ofi: &[f32]) -> Vec<f32> {
use ml_core::state_layout::{MARKET_DIM, MARKET_START, OFI_DIM, OFI_START};
let mut v = vec![0.0_f32; STATE_DIM];
for (i, &m) in market.iter().take(MARKET_DIM).enumerate() {
v[MARKET_START + i] = m;
}
for (i, &o) in ofi.iter().take(OFI_DIM).enumerate() {
v[OFI_START + i] = o;
}
v
}
// ─── Minimal fxcache reader ────────────────────────────────────────────────
//
// The fxcache format is a simple flat binary (documented in ml/src/fxcache.rs):
// Header (64 bytes): magic "FXCACHE\0" + u16 version + u16 feat_dim +
// u16 target_dim + u16 ofi_dim + u64 bar_count +
// [u8; 32] cache_key + [u8; 8] reserved.
// Body (bar_count records): [i64 timestamp][42 f32 features]
// [6 f32 targets][20 f32 ofi] = 280 bytes/bar.
//
// The CLI only needs features + OFI, so this reader skips timestamps + targets.
const FXCACHE_MAGIC: [u8; 8] = *b"FXCACHE\0";
const FXCACHE_HEADER_BYTES: usize = 64;
const FXCACHE_FEAT_DIM: usize = 42;
const FXCACHE_TARGET_DIM: usize = 6;
const FXCACHE_OFI_DIM: usize = 20;
const FXCACHE_EXPECTED_VERSION: u16 = 4;
/// Read an fxcache file and return per-bar (market_features, ofi) pairs.
fn read_fxcache_bars(path: &Path) -> Result<(Vec<Vec<f32>>, Vec<Vec<f32>>)> {
let file = std::fs::File::open(path)
.with_context(|| format!("open fxcache {}", path.display()))?;
let mut reader = BufReader::new(file);
let mut header = [0u8; FXCACHE_HEADER_BYTES];
reader
.read_exact(&mut header)
.context("read fxcache header")?;
let magic: [u8; 8] = header[0..8]
.try_into()
.context("fxcache magic slice")?;
if magic != FXCACHE_MAGIC {
bail!(
"bad fxcache magic: expected {:?} got {:?}",
FXCACHE_MAGIC,
magic
);
}
let version = u16::from_le_bytes(header[8..10].try_into().context("version bytes")?);
if version != FXCACHE_EXPECTED_VERSION {
bail!(
"unsupported fxcache version {version} (expected {FXCACHE_EXPECTED_VERSION}); \
regenerate with `fxt precompute-features`"
);
}
let feat_dim = u16::from_le_bytes(header[10..12].try_into().context("feat_dim bytes")?)
as usize;
let target_dim = u16::from_le_bytes(header[12..14].try_into().context("target_dim bytes")?)
as usize;
let ofi_dim = u16::from_le_bytes(header[14..16].try_into().context("ofi_dim bytes")?)
as usize;
let bar_count = u64::from_le_bytes(header[16..24].try_into().context("bar_count bytes")?)
as usize;
if feat_dim != FXCACHE_FEAT_DIM
|| target_dim != FXCACHE_TARGET_DIM
|| ofi_dim != FXCACHE_OFI_DIM
{
bail!(
"fxcache dim mismatch: feat={feat_dim} target={target_dim} ofi={ofi_dim} \
(expected {FXCACHE_FEAT_DIM}/{FXCACHE_TARGET_DIM}/{FXCACHE_OFI_DIM})"
);
}
if bar_count == 0 {
bail!("fxcache has zero bars");
}
let mut features = Vec::with_capacity(bar_count);
let mut ofi = Vec::with_capacity(bar_count);
// Per-bar layout: i64 ts + 42 f32 feat + 6 f32 targets + 20 f32 ofi.
let mut ts_buf = [0u8; 8];
let mut feat_buf = vec![0u8; feat_dim * 4];
let mut target_buf = vec![0u8; target_dim * 4];
let mut ofi_buf = vec![0u8; ofi_dim * 4];
for bar in 0..bar_count {
reader
.read_exact(&mut ts_buf)
.with_context(|| format!("read timestamp for bar {bar}"))?;
reader
.read_exact(&mut feat_buf)
.with_context(|| format!("read features for bar {bar}"))?;
reader
.read_exact(&mut target_buf)
.with_context(|| format!("read targets for bar {bar}"))?;
reader
.read_exact(&mut ofi_buf)
.with_context(|| format!("read ofi for bar {bar}"))?;
let feat_vec: Vec<f32> = (0..feat_dim)
.map(|i| {
let off = i * 4;
f32::from_le_bytes([
feat_buf[off],
feat_buf[off + 1],
feat_buf[off + 2],
feat_buf[off + 3],
])
})
.collect();
let ofi_vec: Vec<f32> = (0..ofi_dim)
.map(|i| {
let off = i * 4;
f32::from_le_bytes([
ofi_buf[off],
ofi_buf[off + 1],
ofi_buf[off + 2],
ofi_buf[off + 3],
])
})
.collect();
features.push(feat_vec);
ofi.push(ofi_vec);
}
Ok((features, ofi))
}
/// Aggregate per-feature statistics and global completeness for the JSON report.
fn build_report(
feature_names: &[String],
per_feature: &[Vec<f64>],
completeness: &[(f64, f64)],
cli: &Cli,
) -> serde_json::Value {
// Per-feature stats.
let mut features = Vec::with_capacity(feature_names.len());
let mut top_by_mean_abs: Vec<(String, f64, f64, f64)> =
Vec::with_capacity(feature_names.len());
for (i, name) in feature_names.iter().enumerate() {
let values = &per_feature[i];
let mean_signed = values.iter().sum::<f64>() / (values.len().max(1) as f64);
let mean_abs =
values.iter().map(|v| v.abs()).sum::<f64>() / (values.len().max(1) as f64);
let var = values
.iter()
.map(|v| (v - mean_signed).powi(2))
.sum::<f64>()
/ (values.len().max(1) as f64);
let stddev = var.sqrt();
features.push(serde_json::json!({
"name": name,
"mean_abs": mean_abs,
"stddev": stddev,
"mean_signed": mean_signed,
}));
top_by_mean_abs.push((name.clone(), mean_abs, stddev, mean_signed));
}
// Top-10 by mean_abs.
top_by_mean_abs.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let top10: Vec<serde_json::Value> = top_by_mean_abs
.iter()
.take(10)
.map(|(name, mean_abs, stddev, mean_signed)| {
serde_json::json!({
"name": name,
"mean_abs": mean_abs,
"stddev": stddev,
"mean_signed": mean_signed,
})
})
.collect();
// Completeness axiom: check that sum(attributions) ~ F(input) - F(baseline)
// for each state; report the max relative error across all states.
let mut worst_rel_err = 0.0_f64;
let mut completeness_details = Vec::with_capacity(completeness.len());
for (i, (sum_attr, expected_diff)) in completeness.iter().enumerate() {
let rel_err = if expected_diff.abs() > 1e-10 {
(sum_attr - expected_diff).abs() / expected_diff.abs()
} else {
(sum_attr - expected_diff).abs()
};
if rel_err > worst_rel_err {
worst_rel_err = rel_err;
}
completeness_details.push(serde_json::json!({
"state_idx": i,
"sum_attributions": sum_attr,
"f_input_minus_f_baseline": expected_diff,
"relative_error": rel_err,
}));
}
serde_json::json!({
"schema_version": 1,
"checkpoint": cli.checkpoint.display().to_string(),
"num_steps": cli.num_steps,
"state_dim": STATE_DIM,
"num_states": per_feature
.first()
.map_or(0, Vec::len),
"forward_target": "mean(Q[direction, 0..4]) == V(s)",
"regime_head": "trending",
"features": features,
"top_10_by_mean_abs": top10,
"completeness": {
"worst_relative_error": worst_rel_err,
"per_state": completeness_details,
},
})
}

View File

@@ -0,0 +1,186 @@
//! End-to-end integration test for the `ig_diag` CLI.
//!
//! 1. Builds a fresh `DQN` with a small config,
//! 2. Saves its branching weights to a safetensors checkpoint with the
//! same plain-name layout that the single-head loader expects,
//! 3. Writes a small states JSON file,
//! 4. Invokes the compiled `ig_diag` binary via `std::process::Command`,
//! 5. Parses the output JSON and asserts the schema + completeness axiom.
//!
//! Gated behind `#[cfg(feature = "cuda")] #[ignore]` because it requires a
//! CUDA runtime and the binary to be built with the `ig-diag-cli` feature.
//! Run with:
//! ```bash
//! cargo build -p ml-explainability --features ig-diag-cli --bin ig_diag
//! cargo test -p ml-explainability --features ig-diag-cli --test \
//! ig_diag_cli_integration -- --ignored --nocapture
//! ```
#![cfg(all(feature = "cuda", feature = "ig-diag-cli"))]
use std::path::PathBuf;
use std::process::Command;
#[test]
#[ignore = "requires CUDA runtime + compiled ig_diag binary"]
fn test_ig_diag_cli_end_to_end() {
// 1. Build a DQN, save a safetensors checkpoint.
let dir = tempfile::tempdir().expect("tempdir");
let ckpt_path = dir.path().join("ig_diag_test.safetensors");
let states_path = dir.path().join("states.json");
let output_path = dir.path().join("ig_report.json");
save_dqn_checkpoint(&ckpt_path);
// 2. Write a small states JSON file (3 random STATE_DIM=96 vectors).
let state_dim = ml_core::state_layout::STATE_DIM;
let states_json = serde_json::json!({
"states": [
(0..state_dim)
.map(|i| ((i as f32 * 0.0123).sin() as f64) * 0.1)
.collect::<Vec<f64>>(),
(0..state_dim)
.map(|i| ((i as f32 * 0.0456 + 1.0).cos() as f64) * 0.2)
.collect::<Vec<f64>>(),
(0..state_dim)
.map(|i| ((i as f32 * 0.0789 - 0.5).sin() as f64) * 0.3)
.collect::<Vec<f64>>(),
],
});
std::fs::write(&states_path, serde_json::to_string(&states_json).unwrap())
.expect("write states.json");
// 3. Invoke the ig_diag binary.
let binary = binary_path();
assert!(
binary.exists(),
"ig_diag binary not found at {} — run: \
cargo build -p ml-explainability --features ig-diag-cli --bin ig_diag",
binary.display(),
);
let status = Command::new(&binary)
.arg("--checkpoint")
.arg(&ckpt_path)
.arg("--states")
.arg(&states_path)
.arg("--num-steps")
.arg("64") // enough steps for <5% completeness on a non-linear DQN
.arg("--output")
.arg(&output_path)
.status()
.expect("spawn ig_diag");
assert!(status.success(), "ig_diag exited non-zero: {status}");
// 4. Parse + validate the output JSON.
let raw = std::fs::read_to_string(&output_path).expect("read output JSON");
let v: serde_json::Value = serde_json::from_str(&raw).expect("parse output JSON");
// Schema checks.
assert_eq!(v["schema_version"], 1);
assert_eq!(v["state_dim"], state_dim);
assert_eq!(v["num_states"], 3);
assert!(
v["features"].is_array(),
"features should be array",
);
assert_eq!(
v["features"].as_array().unwrap().len(),
state_dim,
"one entry per feature",
);
let feat0 = &v["features"][0];
assert!(feat0["mean_abs"].is_number(), "mean_abs should be numeric");
assert!(feat0["stddev"].is_number(), "stddev should be numeric");
assert!(feat0["mean_signed"].is_number(), "mean_signed should be numeric");
let top10 = v["top_10_by_mean_abs"].as_array().expect("top10 array");
assert!(top10.len() <= 10);
assert!(!top10.is_empty(), "top10 must be non-empty");
// Completeness: worst relative error < 5%.
//
// For IG with num_steps=16 on a DQN forward (non-linear), 5% is a loose
// upper bound (in practice it's much tighter). Tighten if needed.
let worst = v["completeness"]["worst_relative_error"]
.as_f64()
.expect("worst_relative_error numeric");
assert!(
worst < 0.05,
"completeness axiom violated: worst relative error {worst} >= 5%",
);
// Per-state completeness entries must be present.
let per_state = v["completeness"]["per_state"]
.as_array()
.expect("per_state array");
assert_eq!(per_state.len(), 3);
println!("ig_diag CLI integration test OK — worst_rel_err={worst:.6}");
}
/// Return the path to the compiled `ig_diag` binary. CARGO_BIN_EXE_* is
/// injected by cargo for integration tests in the same crate; fall back to
/// target/debug if the env var is missing (e.g., when running manually).
fn binary_path() -> PathBuf {
if let Ok(p) = std::env::var("CARGO_BIN_EXE_ig_diag") {
return PathBuf::from(p);
}
// Fallback: walk up from CARGO_MANIFEST_DIR to find workspace/target/debug.
let mut p = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
while p.pop() {
let candidate = p.join("target").join("debug").join("ig_diag");
if candidate.exists() {
return candidate;
}
}
PathBuf::from("target/debug/ig_diag")
}
/// Build a DQN with small hidden dims, disable NoisyNet, dump its branching
/// weights to a plain-name safetensors file. Mirrors the minimal save path
/// used by `DqnInferenceAdapter`'s round-trip test.
fn save_dqn_checkpoint(path: &std::path::Path) {
use ml_dqn::{DQN, DQNConfig};
let mut config = DQNConfig::emergency_safe_defaults();
config.num_actions = 3;
config.hidden_dims = vec![64, 64];
let dqn = DQN::new(config.clone()).expect("DQN new");
let stream = dqn.cuda_stream().clone();
let branching = dqn
.branching_q_network
.as_ref()
.expect("branching_q_network");
// Download each tensor + serialize with safetensors using the plain-name
// layout ({shared,value,branch}_*). The loader accepts this or the
// regime-prefixed layout.
let named = branching.named_weight_slices();
let byte_vecs: Vec<(String, Vec<usize>, Vec<u8>)> = named
.iter()
.map(|(name, slice, shape)| {
let mut host = vec![0.0_f32; slice.len()];
stream.memcpy_dtoh(*slice, &mut host).expect("DtoH");
let bytes: Vec<u8> = host.iter().flat_map(|v| v.to_le_bytes()).collect();
(name.clone(), shape.clone(), bytes)
})
.collect();
let mut tensor_map = std::collections::HashMap::new();
for (name, shape, bytes) in &byte_vecs {
let view = safetensors::tensor::TensorView::new(
safetensors::Dtype::F32,
shape.clone(),
bytes,
)
.expect("TensorView");
tensor_map.insert(name.as_str(), view);
}
let arch_meta = Some(config.checkpoint_metadata());
safetensors::serialize_to_file(tensor_map, arch_meta, path)
.expect("serialize_to_file");
}