optirs_core/quantum_inspired/
hybrid.rs1use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
10use scirs2_core::numeric::Float;
11use std::fmt::Debug;
12
13use crate::error::{OptimError, Result};
14use crate::optimizers::{Adam, Optimizer};
15
16use super::annealing::QuantumAnnealing;
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum OptimizationPhase {
21 Exploration,
23 Refinement,
25}
26
27#[derive(Debug)]
51pub struct HybridQuantumClassical<A: Float + ScalarOperand + Debug> {
52 quantum: QuantumAnnealing<A>,
54 classical: Adam<A>,
56 phase: OptimizationPhase,
58 switch_step: usize,
60 current_step: usize,
62 warmed_up: bool,
64}
65
66impl<A> HybridQuantumClassical<A>
67where
68 A: Float + ScalarOperand + Debug + Send + Sync,
69{
70 pub fn new(learning_rate: A, switch_step: usize) -> Self {
72 let initial_phase = if switch_step == 0 {
73 OptimizationPhase::Refinement
74 } else {
75 OptimizationPhase::Exploration
76 };
77 Self {
78 quantum: QuantumAnnealing::new(learning_rate),
79 classical: Adam::new(learning_rate),
80 phase: initial_phase,
81 switch_step,
82 current_step: 0,
83 warmed_up: false,
84 }
85 }
86
87 pub fn with_quantum(mut self, quantum: QuantumAnnealing<A>) -> Self {
89 self.quantum = quantum;
90 self
91 }
92
93 pub fn with_classical(mut self, classical: Adam<A>) -> Self {
95 self.classical = classical;
96 self
97 }
98
99 pub fn with_quantum_config<F>(mut self, f: F) -> Self
101 where
102 F: FnOnce(QuantumAnnealing<A>) -> QuantumAnnealing<A>,
103 {
104 self.quantum = f(self.quantum);
105 self
106 }
107
108 pub fn with_classical_config<F>(mut self, f: F) -> Self
110 where
111 F: FnOnce(Adam<A>) -> Adam<A>,
112 {
113 self.classical = f(self.classical);
114 self
115 }
116
117 pub fn phase(&self) -> OptimizationPhase {
119 self.phase
120 }
121
122 pub fn switch_step(&self) -> usize {
124 self.switch_step
125 }
126
127 pub fn current_step(&self) -> usize {
129 self.current_step
130 }
131
132 pub fn warmed_up(&self) -> bool {
134 self.warmed_up
135 }
136
137 pub fn learning_rate(&self) -> A {
140 match self.phase {
141 OptimizationPhase::Exploration => self.quantum.learning_rate(),
142 OptimizationPhase::Refinement => self.classical.learning_rate(),
143 }
144 }
145
146 pub fn set_lr(&mut self, learning_rate: A) {
149 self.quantum.set_lr(learning_rate);
150 self.classical.set_lr(learning_rate);
151 }
152
153 pub fn quantum(&self) -> &QuantumAnnealing<A> {
155 &self.quantum
156 }
157
158 pub fn classical(&self) -> &Adam<A> {
160 &self.classical
161 }
162
163 pub fn quantum_mut(&mut self) -> &mut QuantumAnnealing<A> {
165 &mut self.quantum
166 }
167
168 pub fn classical_mut(&mut self) -> &mut Adam<A> {
170 &mut self.classical
171 }
172
173 fn maybe_warm_start<D>(&mut self, params: &Array<A, D>) -> Result<Array<A, D>>
177 where
178 D: Dimension,
179 {
180 if self.warmed_up {
181 return Ok(params.clone());
182 }
183 let anchor: Array<A, D> = match self.quantum.best_params::<D>() {
184 Some(best) if best.shape() == params.shape() => best,
185 _ => params.clone(),
186 };
187 let zero_grads = Array::<A, D>::zeros(anchor.raw_dim());
191 let warmed = self.classical.step(&anchor, &zero_grads)?;
192 self.warmed_up = true;
193 Ok(warmed)
194 }
195}
196
197impl<A, D> Optimizer<A, D> for HybridQuantumClassical<A>
198where
199 A: Float + ScalarOperand + Debug + Send + Sync,
200 D: Dimension,
201{
202 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
203 if params.shape() != gradients.shape() {
204 return Err(OptimError::DimensionMismatch(format!(
205 "Hybrid optimizer: parameters have shape {:?}, gradients have shape {:?}",
206 params.shape(),
207 gradients.shape()
208 )));
209 }
210
211 let in_exploration = self.current_step < self.switch_step;
212 if in_exploration {
213 self.phase = OptimizationPhase::Exploration;
214 let next = self.quantum.step(params, gradients)?;
215 self.current_step = self.current_step.saturating_add(1);
216 Ok(next)
217 } else {
218 let anchor = self.maybe_warm_start(params)?;
220 self.phase = OptimizationPhase::Refinement;
221 let next = self.classical.step(&anchor, gradients)?;
222 self.current_step = self.current_step.saturating_add(1);
223 Ok(next)
224 }
225 }
226
227 fn get_learning_rate(&self) -> A {
228 match self.phase {
229 OptimizationPhase::Exploration => self.quantum.learning_rate(),
230 OptimizationPhase::Refinement => self.classical.learning_rate(),
231 }
232 }
233
234 fn set_learning_rate(&mut self, learning_rate: A) {
235 self.quantum.set_lr(learning_rate);
236 self.classical.set_lr(learning_rate);
237 }
238}
239
240#[cfg(test)]
241mod tests {
242 use super::*;
243 use approx::assert_abs_diff_eq;
244 use scirs2_core::ndarray::Array1;
245
246 fn quadratic_grad(params: &Array1<f64>) -> Array1<f64> {
247 params.mapv(|x| 2.0 * x)
248 }
249
250 #[test]
251 fn test_starts_in_exploration_phase() {
252 let optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.05, 30);
253 assert_eq!(optimizer.phase(), OptimizationPhase::Exploration);
254 assert_eq!(optimizer.current_step(), 0);
255 assert_eq!(optimizer.switch_step(), 30);
256 assert!(!optimizer.warmed_up());
257 }
258
259 #[test]
260 fn test_switches_at_correct_step() {
261 let mut optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.05, 5)
262 .with_quantum_config(|q| {
263 q.with_temperature_schedule(1.0, 1e-3)
264 .with_iterations(50)
265 .with_seed(11)
266 });
267 let mut params = Array1::from_vec(vec![1.0, -1.0]);
268 for _ in 0..5 {
270 let grads = quadratic_grad(¶ms);
271 params = optimizer.step(¶ms, &grads).expect("step failed");
272 assert_eq!(optimizer.phase(), OptimizationPhase::Exploration);
273 }
274 let grads = quadratic_grad(¶ms);
276 params = optimizer.step(¶ms, &grads).expect("step failed");
277 assert_eq!(optimizer.phase(), OptimizationPhase::Refinement);
278 let _ = params; }
280
281 #[test]
282 fn test_classical_used_after_switch() {
283 let mut optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.1, 3)
284 .with_quantum_config(|q| {
285 q.with_temperature_schedule(1.0, 1e-3)
286 .with_iterations(20)
287 .with_seed(101)
288 });
289 let mut params = Array1::from_vec(vec![2.0]);
290 for _ in 0..3 {
291 let grads = quadratic_grad(¶ms);
292 params = optimizer.step(¶ms, &grads).expect("step failed");
293 }
294 assert!(!optimizer.warmed_up());
295 let grads = quadratic_grad(¶ms);
296 let _ = optimizer.step(¶ms, &grads).expect("step failed");
297 assert!(optimizer.warmed_up());
298 assert_eq!(optimizer.phase(), OptimizationPhase::Refinement);
299 }
300
301 #[test]
302 fn test_phase_transition_preserves_best_params() {
303 let mut optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.05, 30)
304 .with_quantum_config(|q| {
305 q.with_temperature_schedule(1.0, 1e-3)
306 .with_tunneling(0.1)
307 .with_iterations(60)
308 .with_seed(2024)
309 });
310 let mut params = Array1::from_vec(vec![5.0]);
311 for _ in 0..30 {
313 let grads = quadratic_grad(¶ms);
314 params = optimizer.step(¶ms, &grads).expect("step failed");
315 }
316 let best_energy_at_switch = optimizer.quantum().best_energy();
317 let grads = quadratic_grad(¶ms);
319 let _ = optimizer.step(¶ms, &grads).expect("step failed");
320 assert!(
322 optimizer.quantum().best_energy() <= best_energy_at_switch + 1e-12,
323 "best_energy regressed after switch: before={best_energy_at_switch}, after={}",
324 optimizer.quantum().best_energy()
325 );
326 }
327
328 #[test]
329 fn test_convergence_on_quadratic() {
330 let mut optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.05, 50)
332 .with_quantum_config(|q| {
333 q.with_temperature_schedule(1.0, 1e-4)
334 .with_tunneling(0.05)
335 .with_iterations(50)
336 .with_seed(7)
337 });
338 let mut params = Array1::from_vec(vec![5.0]);
339 for _ in 0..400 {
340 let grads = quadratic_grad(¶ms);
341 params = optimizer.step(¶ms, &grads).expect("step failed");
342 }
343 assert!(
344 params[0].abs() < 0.5,
345 "Hybrid did not converge on quadratic: |x|={}",
346 params[0].abs()
347 );
348 }
349
350 #[test]
351 fn test_switch_step_zero_means_pure_classical() {
352 let mut optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.1, 0);
353 assert_eq!(optimizer.phase(), OptimizationPhase::Refinement);
355 let params = Array1::from_vec(vec![2.0]);
356 let grads = quadratic_grad(¶ms);
357 let _ = optimizer.step(¶ms, &grads).expect("step failed");
358 assert_eq!(optimizer.phase(), OptimizationPhase::Refinement);
359 assert!(optimizer.warmed_up());
360 }
361
362 #[test]
363 fn test_switch_step_inf_means_pure_quantum() {
364 let mut optimizer: HybridQuantumClassical<f64> =
365 HybridQuantumClassical::new(0.05, usize::MAX).with_quantum_config(|q| {
366 q.with_temperature_schedule(1.0, 1e-3)
367 .with_iterations(100)
368 .with_seed(5)
369 });
370 let mut params = Array1::from_vec(vec![1.0, -1.0]);
371 for _ in 0..50 {
372 let grads = quadratic_grad(¶ms);
373 params = optimizer.step(¶ms, &grads).expect("step failed");
374 assert_eq!(optimizer.phase(), OptimizationPhase::Exploration);
375 }
376 assert!(!optimizer.warmed_up());
377 }
378
379 #[test]
380 fn test_builder_pattern() {
381 let optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.05, 20)
382 .with_quantum_config(|q| {
383 q.with_temperature_schedule(1.5, 1e-2)
384 .with_tunneling(0.4)
385 .with_iterations(100)
386 .with_seed(42)
387 })
388 .with_classical_config(|adam| adam.with_beta1(0.95).with_beta2(0.9999));
389 assert_eq!(optimizer.switch_step(), 20);
390 assert_abs_diff_eq!(optimizer.quantum().initial_temperature(), 1.5);
391 assert_abs_diff_eq!(optimizer.quantum().final_temperature(), 1e-2);
392 assert_abs_diff_eq!(optimizer.quantum().tunneling_strength(), 0.4);
393 assert_eq!(optimizer.quantum().num_iterations(), 100);
394 assert_eq!(optimizer.quantum().seed(), 42);
395 assert_abs_diff_eq!(optimizer.classical().get_beta1(), 0.95);
396 assert_abs_diff_eq!(optimizer.classical().get_beta2(), 0.9999);
397 }
398
399 #[test]
400 fn test_dimension_mismatch_errors() {
401 let mut optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.05, 10);
402 let params = Array1::from_vec(vec![1.0, 2.0]);
403 let grads = Array1::from_vec(vec![0.1]);
404 let result = optimizer.step(¶ms, &grads);
405 assert!(result.is_err(), "expected dimension mismatch error");
406 }
407
408 #[test]
409 fn test_set_learning_rate_applies_to_both() {
410 let mut optimizer: HybridQuantumClassical<f64> = HybridQuantumClassical::new(0.05, 5);
411 optimizer.set_lr(0.2);
412 assert_abs_diff_eq!(optimizer.quantum().learning_rate(), 0.2);
413 assert_abs_diff_eq!(optimizer.classical().learning_rate(), 0.2);
414 }
415}