Skip to main content

optirs_core/quantum_inspired/
hybrid.rs

1// Hybrid quantum-classical optimizer.
2//
3// `HybridQuantumClassical` wraps a quantum-inspired exploration phase
4// (`QuantumAnnealing`) with a classical refinement phase (`Adam`). The
5// optimizer transitions from exploration to refinement after a configurable
6// number of steps, optionally seeding the classical optimizer with the best
7// parameters discovered during exploration.
8
9use 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/// Phase of the hybrid quantum-classical optimization process.
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum OptimizationPhase {
21    /// Exploration via quantum-inspired annealing.
22    Exploration,
23    /// Refinement via classical (Adam) updates.
24    Refinement,
25}
26
27/// Hybrid quantum-classical optimizer.
28///
29/// `HybridQuantumClassical` blends quantum-inspired exploration with classical
30/// refinement, much in the spirit of variational hybrid algorithms. The first
31/// `switch_step` updates are produced by a [`QuantumAnnealing`] instance that
32/// performs broad exploration of the loss landscape. Subsequent updates are
33/// produced by an [`Adam`] instance whose state is implicitly warm-started by
34/// the best parameters discovered by the annealer.
35///
36/// # Examples
37///
38/// ```
39/// use optirs_core::quantum_inspired::HybridQuantumClassical;
40/// use optirs_core::optimizers::Optimizer;
41/// use scirs2_core::ndarray::Array1;
42///
43/// let mut optimizer: HybridQuantumClassical<f64> =
44///     HybridQuantumClassical::new(0.05, 50);
45///
46/// let params = Array1::from_vec(vec![1.0, -1.0, 0.5]);
47/// let grads = Array1::from_vec(vec![0.1, -0.2, 0.05]);
48/// let _ = optimizer.step(&params, &grads).expect("step failed");
49/// ```
50#[derive(Debug)]
51pub struct HybridQuantumClassical<A: Float + ScalarOperand + Debug> {
52    /// Quantum-inspired exploration optimizer.
53    quantum: QuantumAnnealing<A>,
54    /// Classical refinement optimizer.
55    classical: Adam<A>,
56    /// Current phase.
57    phase: OptimizationPhase,
58    /// Step at which to transition from exploration to refinement.
59    switch_step: usize,
60    /// Current step counter.
61    current_step: usize,
62    /// Whether we have performed the warm-start transition yet.
63    warmed_up: bool,
64}
65
66impl<A> HybridQuantumClassical<A>
67where
68    A: Float + ScalarOperand + Debug + Send + Sync,
69{
70    /// Create a new hybrid optimizer with default sub-optimizer settings.
71    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    /// Replace the underlying quantum annealer.
88    pub fn with_quantum(mut self, quantum: QuantumAnnealing<A>) -> Self {
89        self.quantum = quantum;
90        self
91    }
92
93    /// Replace the underlying classical optimizer.
94    pub fn with_classical(mut self, classical: Adam<A>) -> Self {
95        self.classical = classical;
96        self
97    }
98
99    /// Configure the quantum annealer via a builder closure.
100    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    /// Configure the classical optimizer via a builder closure.
109    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    /// Returns the current optimization phase.
118    pub fn phase(&self) -> OptimizationPhase {
119        self.phase
120    }
121
122    /// Returns the configured switch step.
123    pub fn switch_step(&self) -> usize {
124        self.switch_step
125    }
126
127    /// Returns the current step counter.
128    pub fn current_step(&self) -> usize {
129        self.current_step
130    }
131
132    /// Returns whether the warm-start handoff has fired.
133    pub fn warmed_up(&self) -> bool {
134        self.warmed_up
135    }
136
137    /// Returns the current learning rate, choosing the relevant sub-optimizer
138    /// based on the current phase.
139    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    /// Set the learning rate on both the quantum annealer and the classical
147    /// optimizer simultaneously.
148    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    /// Returns an immutable reference to the underlying quantum annealer.
154    pub fn quantum(&self) -> &QuantumAnnealing<A> {
155        &self.quantum
156    }
157
158    /// Returns an immutable reference to the underlying classical optimizer.
159    pub fn classical(&self) -> &Adam<A> {
160        &self.classical
161    }
162
163    /// Returns a mutable reference to the underlying quantum annealer.
164    pub fn quantum_mut(&mut self) -> &mut QuantumAnnealing<A> {
165        &mut self.quantum
166    }
167
168    /// Returns a mutable reference to the underlying classical optimizer.
169    pub fn classical_mut(&mut self) -> &mut Adam<A> {
170        &mut self.classical
171    }
172
173    /// Internal: pick the best params recorded by the annealer (if any) and
174    /// warm-start the classical optimizer by stepping it with zero gradient.
175    /// Returns the params the classical optimizer should treat as its anchor.
176    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        // Step Adam with zero gradient to register the anchor shape in its
188        // internal state without actually changing the value (Adam with grad=0
189        // and no weight decay simply leaves params unchanged).
190        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            // Trigger warm-start transition on the first refinement step.
219            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        // First 5 steps must be exploration.
269        for _ in 0..5 {
270            let grads = quadratic_grad(&params);
271            params = optimizer.step(&params, &grads).expect("step failed");
272            assert_eq!(optimizer.phase(), OptimizationPhase::Exploration);
273        }
274        // Sixth step must be refinement.
275        let grads = quadratic_grad(&params);
276        params = optimizer.step(&params, &grads).expect("step failed");
277        assert_eq!(optimizer.phase(), OptimizationPhase::Refinement);
278        let _ = params; // suppress unused warning if any
279    }
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(&params);
292            params = optimizer.step(&params, &grads).expect("step failed");
293        }
294        assert!(!optimizer.warmed_up());
295        let grads = quadratic_grad(&params);
296        let _ = optimizer.step(&params, &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        // Run exploration phase.
312        for _ in 0..30 {
313            let grads = quadratic_grad(&params);
314            params = optimizer.step(&params, &grads).expect("step failed");
315        }
316        let best_energy_at_switch = optimizer.quantum().best_energy();
317        // One refinement step.
318        let grads = quadratic_grad(&params);
319        let _ = optimizer.step(&params, &grads).expect("step failed");
320        // best_energy is non-increasing because we only ever record improvements.
321        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        // Hybrid should converge at least as well as Adam on a convex bowl.
331        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(&params);
341            params = optimizer.step(&params, &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        // From step 0 we should already be in refinement.
354        assert_eq!(optimizer.phase(), OptimizationPhase::Refinement);
355        let params = Array1::from_vec(vec![2.0]);
356        let grads = quadratic_grad(&params);
357        let _ = optimizer.step(&params, &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(&params);
373            params = optimizer.step(&params, &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(&params, &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}