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:
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user