test(ensemble): add checkpoint round-trip and full integration tests

- DQN adapter: save/load/predict round-trip validation
- Coordinator: full ensemble with real DQN adapter integration test

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-21 13:42:32 +01:00
parent c300fa0551
commit 89079a72c9
2 changed files with 90 additions and 0 deletions

View File

@@ -202,4 +202,56 @@ mod tests {
"Deterministic predictions should have same direction"
);
}
#[test]
fn test_dqn_checkpoint_round_trip() {
let config = test_config();
let adapter = DqnInferenceAdapter::new(config.clone());
assert!(
adapter.is_ok(),
"failed to create adapter: {:?}",
adapter.err()
);
let adapter = adapter.unwrap();
let fv = FeatureVector {
values: vec![0.3; 51],
timestamp: 1700000000,
};
let pred1 = adapter.predict(&fv);
assert!(pred1.is_ok(), "first predict failed: {:?}", pred1.err());
let pred1 = pred1.unwrap();
// Save checkpoint to temp dir
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("dqn_test.safetensors");
{
let model = adapter.model.lock().unwrap();
model.get_q_network_vars().save(path.as_path()).unwrap();
}
// Load into new adapter
let path_str = path.to_str().unwrap();
let adapter2 = DqnInferenceAdapter::from_checkpoint(config, path_str);
assert!(
adapter2.is_ok(),
"from_checkpoint failed: {:?}",
adapter2.err()
);
let adapter2 = adapter2.unwrap();
let pred2 = adapter2.predict(&fv);
assert!(
pred2.is_ok(),
"second predict failed: {:?}",
pred2.err()
);
let pred2 = pred2.unwrap();
assert!(
(pred1.direction - pred2.direction).abs() < 1e-4,
"round-trip mismatch: {} vs {}",
pred1.direction,
pred2.direction
);
}
}

View File

@@ -761,6 +761,44 @@ mod tests {
assert_eq!(decision.model_count(), 1);
}
#[tokio::test]
async fn test_full_ensemble_with_dqn_adapter() {
use crate::ensemble::adapters::DqnInferenceAdapter;
use crate::dqn::dqn::DQNConfig;
let dqn_config = DQNConfig {
state_dim: 51,
num_actions: 45,
hidden_dims: vec![64, 64],
..Default::default()
};
let adapter = DqnInferenceAdapter::new(dqn_config).unwrap();
let mut coordinator = EnsembleCoordinator::new();
coordinator.add_adapter(Box::new(adapter));
coordinator
.register_model("DQN".to_string(), 1.0)
.await
.unwrap();
let features = Features::new(
vec![
0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0, 0.1, 0.2, 0.3, 0.4, 0.5,
0.6, 0.7, 0.8, 0.9, 1.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0,
0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0, 0.1, 0.2, 0.3, 0.4, 0.5,
0.6, 0.7, 0.8, 0.9, 1.0, 0.55,
],
(0..51).map(|i| format!("f{}", i)).collect(),
);
let decision = coordinator.predict(&features).await.unwrap();
assert!(decision.confidence > 0.0);
assert!(decision.signal >= -1.0 && decision.signal <= 1.0);
assert_eq!(decision.model_count(), 1);
}
#[tokio::test]
async fn test_ensemble_graceful_no_adapters() {
let coordinator = EnsembleCoordinator::new();