Skip to main content

quantrs2_anneal/active_learning_decomposition/
strategy_learning.rs

1//! Strategy learning components for active learning decomposition
2
3use scirs2_core::ndarray::Array1;
4use std::collections::HashMap;
5use std::time::{Duration, Instant};
6
7use super::{
8    DecompositionStrategy, DiversityMetric, DomainAdaptationStrategy, EvaluationMetric, ModelType,
9    ProblemAnalysis, QueryStrategy, StructureType,
10};
11use crate::ising::IsingModel;
12
13/// Decomposition strategy learner
14#[derive(Debug, Clone)]
15pub struct DecompositionStrategyLearner {
16    /// Strategy selection model
17    pub selection_model: StrategySelectionModel,
18    /// Strategy performance history
19    pub performance_history: HashMap<String, Vec<PerformanceRecord>>,
20    /// Active learning query selector
21    pub query_selector: QuerySelector,
22    /// Transfer learning manager
23    pub transfer_learning: TransferLearningManager,
24    /// Learning statistics
25    pub learning_stats: LearningStatistics,
26}
27
28impl DecompositionStrategyLearner {
29    pub fn new() -> Result<Self, String> {
30        Ok(Self {
31            selection_model: StrategySelectionModel::new(),
32            performance_history: HashMap::new(),
33            query_selector: QuerySelector::new(),
34            transfer_learning: TransferLearningManager::new(),
35            learning_stats: LearningStatistics::new(),
36        })
37    }
38
39    pub const fn recommend_strategy(
40        &self,
41        problem: &IsingModel,
42        analysis: &ProblemAnalysis,
43    ) -> Result<DecompositionStrategy, String> {
44        // Simplified strategy recommendation based on problem size
45        if problem.num_qubits < 10 {
46            Ok(DecompositionStrategy::NoDecomposition)
47        } else if problem.num_qubits < 50 {
48            Ok(DecompositionStrategy::GraphPartitioning)
49        } else {
50            Ok(DecompositionStrategy::CommunityDetection)
51        }
52    }
53}
54
55/// Strategy selection model
56#[derive(Debug, Clone)]
57pub struct StrategySelectionModel {
58    /// Model type
59    pub model_type: ModelType,
60    /// Feature weights
61    pub feature_weights: Array1<f64>,
62    /// Strategy preferences
63    pub strategy_preferences: HashMap<DecompositionStrategy, f64>,
64    /// Uncertainty estimates
65    pub uncertainty_estimates: HashMap<String, f64>,
66    /// Model parameters
67    pub model_parameters: ModelParameters,
68}
69
70impl StrategySelectionModel {
71    #[must_use]
72    pub fn new() -> Self {
73        Self {
74            model_type: ModelType::Linear,
75            feature_weights: Array1::ones(20),
76            strategy_preferences: HashMap::new(),
77            uncertainty_estimates: HashMap::new(),
78            model_parameters: ModelParameters::default(),
79        }
80    }
81
82    pub fn get_uncertainty(&self, features: &Array1<f64>) -> Result<f64, String> {
83        // Simplified uncertainty calculation
84        let feature_sum = features.sum();
85        Ok(1.0 / (1.0 + feature_sum.abs()))
86    }
87
88    pub fn get_strategy_uncertainty(
89        &self,
90        strategy: &DecompositionStrategy,
91        features: &Array1<f64>,
92    ) -> Result<f64, String> {
93        let base_uncertainty = self.get_uncertainty(features)?;
94        let strategy_key = format!("{strategy:?}");
95
96        if let Some(&stored_uncertainty) = self.uncertainty_estimates.get(&strategy_key) {
97            Ok(f64::midpoint(base_uncertainty, stored_uncertainty))
98        } else {
99            Ok(base_uncertainty)
100        }
101    }
102}
103
104/// Model parameters
105#[derive(Debug, Clone)]
106pub struct ModelParameters {
107    /// Model-specific parameters
108    pub parameters: HashMap<String, f64>,
109    /// Regularization parameters
110    pub regularization: RegularizationParameters,
111    /// Training configuration
112    pub training_config: ModelTrainingConfig,
113}
114
115impl Default for ModelParameters {
116    fn default() -> Self {
117        Self {
118            parameters: HashMap::new(),
119            regularization: RegularizationParameters {
120                l1_weight: 0.01,
121                l2_weight: 0.01,
122                dropout_rate: 0.1,
123                early_stopping_patience: 10,
124            },
125            training_config: ModelTrainingConfig {
126                num_epochs: 100,
127                batch_size: 32,
128                learning_rate: 0.001,
129                validation_split: 0.2,
130            },
131        }
132    }
133}
134
135/// Regularization parameters
136#[derive(Debug, Clone)]
137pub struct RegularizationParameters {
138    /// L1 regularization weight
139    pub l1_weight: f64,
140    /// L2 regularization weight
141    pub l2_weight: f64,
142    /// Dropout rate
143    pub dropout_rate: f64,
144    /// Early stopping patience
145    pub early_stopping_patience: usize,
146}
147
148/// Model training configuration
149#[derive(Debug, Clone)]
150pub struct ModelTrainingConfig {
151    /// Number of training epochs
152    pub num_epochs: usize,
153    /// Batch size
154    pub batch_size: usize,
155    /// Learning rate
156    pub learning_rate: f64,
157    /// Validation split
158    pub validation_split: f64,
159}
160
161/// Query selector for active learning
162#[derive(Debug, Clone)]
163pub struct QuerySelector {
164    /// Query strategy
165    pub query_strategy: QueryStrategy,
166    /// Uncertainty threshold
167    pub uncertainty_threshold: f64,
168    /// Diversity constraint
169    pub diversity_constraint: DiversityConstraint,
170    /// Query history
171    pub query_history: Vec<QueryRecord>,
172}
173
174impl QuerySelector {
175    #[must_use]
176    pub const fn new() -> Self {
177        Self {
178            query_strategy: QueryStrategy::UncertaintySampling,
179            uncertainty_threshold: 0.5,
180            diversity_constraint: DiversityConstraint {
181                min_distance: 0.1,
182                diversity_metric: DiversityMetric::Euclidean,
183                max_similarity: 0.8,
184            },
185            query_history: Vec::new(),
186        }
187    }
188}
189
190/// Diversity constraint for query selection
191#[derive(Debug, Clone)]
192pub struct DiversityConstraint {
193    /// Minimum distance between queries
194    pub min_distance: f64,
195    /// Diversity metric
196    pub diversity_metric: DiversityMetric,
197    /// Maximum similarity allowed
198    pub max_similarity: f64,
199}
200
201/// Query record
202#[derive(Debug, Clone)]
203pub struct QueryRecord {
204    /// Query timestamp
205    pub timestamp: Instant,
206    /// Queried problem features
207    pub problem_features: Array1<f64>,
208    /// Recommended strategy
209    pub recommended_strategy: DecompositionStrategy,
210    /// Query outcome
211    pub query_outcome: QueryOutcome,
212    /// Performance feedback
213    pub performance_feedback: Option<PerformanceRecord>,
214}
215
216/// Query outcome
217#[derive(Debug, Clone)]
218pub struct QueryOutcome {
219    /// Strategy actually used
220    pub strategy_used: DecompositionStrategy,
221    /// User accepted recommendation
222    pub accepted_recommendation: bool,
223    /// Performance achieved
224    pub performance_achieved: f64,
225    /// Feedback quality
226    pub feedback_quality: f64,
227}
228
229/// Transfer learning manager
230#[derive(Debug, Clone)]
231pub struct TransferLearningManager {
232    /// Source domain models
233    pub source_models: Vec<SourceDomainModel>,
234    /// Domain adaptation strategy
235    pub adaptation_strategy: DomainAdaptationStrategy,
236    /// Knowledge transfer weights
237    pub transfer_weights: Array1<f64>,
238    /// Transfer learning statistics
239    pub transfer_stats: TransferStatistics,
240}
241
242impl TransferLearningManager {
243    #[must_use]
244    pub fn new() -> Self {
245        Self {
246            source_models: Vec::new(),
247            adaptation_strategy: DomainAdaptationStrategy::FineTuning,
248            transfer_weights: Array1::ones(5),
249            transfer_stats: TransferStatistics {
250                successful_transfers: 0,
251                failed_transfers: 0,
252                avg_transfer_benefit: 0.0,
253                transfer_time_overhead: Duration::from_secs(0),
254            },
255        }
256    }
257}
258
259/// Source domain model
260#[derive(Debug, Clone)]
261pub struct SourceDomainModel {
262    /// Domain identifier
263    pub domain_id: String,
264    /// Model for this domain
265    pub model: StrategySelectionModel,
266    /// Domain characteristics
267    pub domain_characteristics: DomainCharacteristics,
268    /// Transfer applicability score
269    pub applicability_score: f64,
270}
271
272/// Domain characteristics
273#[derive(Debug, Clone)]
274pub struct DomainCharacteristics {
275    /// Problem types in domain
276    pub problem_types: Vec<String>,
277    /// Average problem size
278    pub avg_problem_size: f64,
279    /// Problem complexity distribution
280    pub complexity_distribution: Array1<f64>,
281    /// Common structures
282    pub common_structures: Vec<StructureType>,
283}
284
285/// Transfer learning statistics
286#[derive(Debug, Clone)]
287pub struct TransferStatistics {
288    /// Successful transfers
289    pub successful_transfers: usize,
290    /// Failed transfers
291    pub failed_transfers: usize,
292    /// Average transfer benefit
293    pub avg_transfer_benefit: f64,
294    /// Transfer time overhead
295    pub transfer_time_overhead: Duration,
296}
297
298/// Learning statistics
299#[derive(Debug, Clone)]
300pub struct LearningStatistics {
301    /// Total queries made
302    pub total_queries: usize,
303    /// Successful predictions
304    pub successful_predictions: usize,
305    /// Average prediction accuracy
306    pub avg_prediction_accuracy: f64,
307    /// Learning curve data
308    pub learning_curve: Vec<(usize, f64)>, // (query_count, accuracy)
309    /// Exploration vs exploitation ratio
310    pub exploration_exploitation_ratio: f64,
311}
312
313impl LearningStatistics {
314    #[must_use]
315    pub const fn new() -> Self {
316        Self {
317            total_queries: 0,
318            successful_predictions: 0,
319            avg_prediction_accuracy: 0.0,
320            learning_curve: Vec::new(),
321            exploration_exploitation_ratio: 0.5,
322        }
323    }
324}
325
326/// Performance record
327#[derive(Debug, Clone)]
328pub struct PerformanceRecord {
329    /// Timestamp
330    pub timestamp: Instant,
331    /// Problem identifier
332    pub problem_id: String,
333    /// Strategy used
334    pub strategy_used: DecompositionStrategy,
335    /// Performance metrics
336    pub metrics: HashMap<EvaluationMetric, f64>,
337    /// Overall performance score
338    pub overall_score: f64,
339}