diff --git a/ml/src/ensemble/adapters/dqn.rs b/ml/src/ensemble/adapters/dqn.rs index cac1d993e..7653d85be 100644 --- a/ml/src/ensemble/adapters/dqn.rs +++ b/ml/src/ensemble/adapters/dqn.rs @@ -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 + ); + } } diff --git a/ml/src/ensemble/coordinator.rs b/ml/src/ensemble/coordinator.rs index 2511ecbc0..04ae84c6a 100644 --- a/ml/src/ensemble/coordinator.rs +++ b/ml/src/ensemble/coordinator.rs @@ -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();