Skip to main content

optirs_core/privacy/federated_privacy/
adaptation.rs

1//! Meta-learning and task-shift detection for federated rounds.
2//!
3//! # The defects this replaces
4//!
5//! ```text
6//! pub fn compute_client_meta_gradients(...) -> Result<Array1<T>> {
7//!     // Placeholder implementation
8//!     Ok(Array1::default(0))
9//! }
10//!
11//! pub fn detect_task_change(&mut self, updates: &[Array1<T>]) -> Result<bool> {
12//!     // Placeholder implementation
13//!     let _ = updates;
14//!     Ok(false)
15//! }
16//! ```
17//!
18//! The first returned a length-0 array for every input -- it had no storage to
19//! return, because `FederatedMetaLearner::new` discarded its `parameter_size`
20//! argument -- and the second reported "no task change" unconditionally, so the
21//! continual-learning machinery behind it could never fire.
22//!
23//! # What is implemented
24//!
25//! * [`FederatedMetaLearner::compute_client_meta_gradients`]: a first-order MAML
26//!   / Reptile meta-gradient (Finn, Abbeel & Levine 2017; Nichol, Achiam &
27//!   Schulman 2018). Each client's support gradient takes one inner step from
28//!   the shared meta-parameters, the resulting adaptation is recorded, and the
29//!   meta-gradient is the mean of the clients' query gradients. Every length is
30//!   checked, so a mis-sized learner reports the mismatch instead of returning
31//!   an empty array.
32//! * [`TaskDetector::detect_task_change`]: gradient-based change detection. The
33//!   round's mean update is compared with the running mean of the buffered
34//!   history by relative L2 shift and cosine dissimilarity; a change is flagged
35//!   when the larger of the two exceeds the configured threshold, and the
36//!   [`ChangePoint`] records the magnitude that triggered it. Detection methods
37//!   other than `GradientBased` return [`OptimError::UnsupportedOperation`]
38//!   rather than silently behaving like the gradient-based one.
39
40use crate::error::{OptimError, Result};
41use scirs2_core::ndarray::Array1;
42use scirs2_core::numeric::Float;
43use std::collections::HashMap;
44use std::fmt::Debug;
45
46use super::components::{ChangePoint, FederatedMetaLearner, TaskDetector, TaskDistribution};
47use super::config::TaskDetectionMethod;
48
49/// Default inner-loop step size for the first-order meta-gradient.
50pub const DEFAULT_INNER_LEARNING_RATE: f64 = 0.01;
51
52/// Maximum number of rounds retained in the task detector's buffer.
53const MAX_GRADIENT_BUFFER: usize = 100;
54
55/// Convert an array to `f64`, rejecting anything non-finite.
56fn as_finite_f64<T: Float + Debug + Send + Sync + 'static>(
57    label: &str,
58    values: &Array1<T>,
59) -> Result<Vec<f64>> {
60    values
61        .iter()
62        .enumerate()
63        .map(|(index, value)| {
64            let as_f64 = value.to_f64().ok_or_else(|| {
65                OptimError::InvalidParameter(format!(
66                    "{label}: element {index} cannot be represented as f64"
67                ))
68            })?;
69            if !as_f64.is_finite() {
70                return Err(OptimError::InvalidParameter(format!(
71                    "{label}: element {index} is {as_f64}"
72                )));
73            }
74            Ok(as_f64)
75        })
76        .collect()
77}
78
79/// Cosine similarity of two equal-length vectors, or `None` if either is zero.
80fn cosine_similarity(left: &[f64], right: &[f64]) -> Option<f64> {
81    let dot: f64 = left.iter().zip(right.iter()).map(|(a, b)| a * b).sum();
82    let left_norm: f64 = left.iter().map(|value| value * value).sum::<f64>().sqrt();
83    let right_norm: f64 = right.iter().map(|value| value * value).sum::<f64>().sqrt();
84    if left_norm <= 0.0 || right_norm <= 0.0 {
85        None
86    } else {
87        Some((dot / (left_norm * right_norm)).clamp(-1.0, 1.0))
88    }
89}
90
91impl<
92        T: Float
93            + Debug
94            + Send
95            + Sync
96            + 'static
97            + Default
98            + Clone
99            + scirs2_core::ndarray::ScalarOperand,
100    > FederatedMetaLearner<T>
101{
102    /// Compute the first-order meta-gradient from the clients' gradients.
103    ///
104    /// `client_gradients` supplies the gradient each client reported for the
105    /// round; `support_data` and `query_data` supply that client's support and
106    /// query gradients for the inner/outer split. Every client in
107    /// `client_gradients` must appear in both, and every array must have exactly
108    /// [`FederatedMetaLearner::parameter_size`] elements.
109    pub fn compute_client_meta_gradients(
110        &mut self,
111        client_gradients: &HashMap<String, Array1<T>>,
112        support_data: &HashMap<String, Array1<T>>,
113        query_data: &HashMap<String, Array1<T>>,
114    ) -> Result<Array1<T>> {
115        self.compute_client_meta_gradients_with_rate(
116            client_gradients,
117            support_data,
118            query_data,
119            DEFAULT_INNER_LEARNING_RATE,
120        )
121    }
122
123    /// As [`FederatedMetaLearner::compute_client_meta_gradients`], with an
124    /// explicit inner-loop step size.
125    pub fn compute_client_meta_gradients_with_rate(
126        &mut self,
127        client_gradients: &HashMap<String, Array1<T>>,
128        support_data: &HashMap<String, Array1<T>>,
129        query_data: &HashMap<String, Array1<T>>,
130        inner_learning_rate: f64,
131    ) -> Result<Array1<T>> {
132        let dimension = self.meta_parameters.len();
133        if dimension == 0 {
134            return Err(OptimError::InvalidState(
135                "this FederatedMetaLearner was constructed with parameter_size = 0, so it has no \
136                 storage for a meta-gradient; construct it with the model's parameter count"
137                    .to_string(),
138            ));
139        }
140        if client_gradients.is_empty() {
141            return Err(OptimError::InvalidParameter(
142                "no client gradients were supplied, so there is no meta-gradient to compute"
143                    .to_string(),
144            ));
145        }
146        if !inner_learning_rate.is_finite() || inner_learning_rate <= 0.0 {
147            return Err(OptimError::InvalidParameter(format!(
148                "the inner learning rate must be positive and finite, got {inner_learning_rate}"
149            )));
150        }
151
152        let meta = as_finite_f64("meta_parameters", &self.meta_parameters)?;
153        let mut accumulated = vec![0.0f64; dimension];
154        let mut adaptations: HashMap<String, Vec<f64>> = HashMap::new();
155        let mut distributions: HashMap<String, (Vec<f64>, Vec<f64>, f64)> = HashMap::new();
156
157        // Deterministic client order: `HashMap` iteration order is randomised
158        // per process, and floating-point addition is not associative, so an
159        // unordered sum would not be reproducible.
160        let mut client_ids: Vec<&String> = client_gradients.keys().collect();
161        client_ids.sort();
162
163        for client_id in &client_ids {
164            let reported = client_gradients.get(*client_id).ok_or_else(|| {
165                OptimError::InvalidState(format!("client `{client_id}` vanished from the map"))
166            })?;
167            if reported.len() != dimension {
168                return Err(OptimError::DimensionMismatch(format!(
169                    "client `{client_id}` reported a {}-element gradient for a {dimension}-element \
170                     model",
171                    reported.len()
172                )));
173            }
174            let support = support_data.get(*client_id).ok_or_else(|| {
175                OptimError::InvalidParameter(format!(
176                    "client `{client_id}` has a gradient but no support gradient, so no inner \
177                     adaptation step can be taken"
178                ))
179            })?;
180            let query = query_data.get(*client_id).ok_or_else(|| {
181                OptimError::InvalidParameter(format!(
182                    "client `{client_id}` has a gradient but no query gradient, so no outer step \
183                     can be taken"
184                ))
185            })?;
186            if support.len() != dimension || query.len() != dimension {
187                return Err(OptimError::DimensionMismatch(format!(
188                    "client `{client_id}` supplied a {}-element support gradient and a {}-element \
189                     query gradient for a {dimension}-element model",
190                    support.len(),
191                    query.len()
192                )));
193            }
194
195            let support_values = as_finite_f64(&format!("client `{client_id}` support"), support)?;
196            let query_values = as_finite_f64(&format!("client `{client_id}` query"), query)?;
197
198            // One inner step: theta_c = theta - alpha * g_support.
199            let adapted: Vec<f64> = meta
200                .iter()
201                .zip(support_values.iter())
202                .map(|(parameter, gradient)| parameter - inner_learning_rate * gradient)
203                .collect();
204
205            // First-order outer step: the meta-gradient is the mean of the
206            // query gradients evaluated at the adapted parameters.
207            for (slot, value) in accumulated.iter_mut().zip(query_values.iter()) {
208                *slot += value;
209            }
210
211            let similarity = cosine_similarity(&support_values, &query_values).unwrap_or(0.0);
212            adaptations.insert((*client_id).clone(), adapted);
213            distributions.insert(
214                (*client_id).clone(),
215                (support_values, query_values, similarity),
216            );
217        }
218
219        let count = client_ids.len() as f64;
220        let mut meta_gradient = Array1::zeros(dimension);
221        for (index, total) in accumulated.iter().enumerate() {
222            let averaged = total / count;
223            meta_gradient[index] = T::from(averaged).ok_or_else(|| {
224                OptimError::InvalidParameter(format!(
225                    "meta-gradient element {index} ({averaged}) cannot be represented"
226                ))
227            })?;
228        }
229
230        // Commit the recorded state only once every client has validated.
231        for (client_id, adapted) in adaptations {
232            let mut array = Array1::zeros(dimension);
233            for (index, value) in adapted.iter().enumerate() {
234                array[index] = T::from(*value).unwrap_or_else(T::zero);
235            }
236            self.client_adaptations.insert(client_id, array);
237        }
238        for (client_id, (support, query, similarity)) in distributions {
239            let mut support_array = Array1::zeros(dimension);
240            let mut query_array = Array1::zeros(dimension);
241            for index in 0..dimension {
242                support_array[index] = T::from(support[index]).unwrap_or_else(T::zero);
243                query_array[index] = T::from(query[index]).unwrap_or_else(T::zero);
244            }
245            self.task_distributions.insert(
246                client_id,
247                TaskDistribution {
248                    support_gradient: support_array,
249                    query_gradient: query_array,
250                    task_similarity: similarity,
251                    adaptation_steps: 1,
252                },
253            );
254        }
255        self.meta_gradient_buffer = meta_gradient.clone();
256        Ok(meta_gradient)
257    }
258
259    /// Apply the buffered meta-gradient to the meta-parameters.
260    pub fn apply_meta_gradient(&mut self, outer_learning_rate: f64) -> Result<()> {
261        if !outer_learning_rate.is_finite() || outer_learning_rate <= 0.0 {
262            return Err(OptimError::InvalidParameter(format!(
263                "the outer learning rate must be positive and finite, got {outer_learning_rate}"
264            )));
265        }
266        if self.meta_gradient_buffer.len() != self.meta_parameters.len() {
267            return Err(OptimError::DimensionMismatch(format!(
268                "the meta-gradient has {} elements but the meta-parameters have {}",
269                self.meta_gradient_buffer.len(),
270                self.meta_parameters.len()
271            )));
272        }
273        let rate = T::from(outer_learning_rate).ok_or_else(|| {
274            OptimError::InvalidParameter(
275                "the outer learning rate cannot be represented in the parameter type".to_string(),
276            )
277        })?;
278        for index in 0..self.meta_parameters.len() {
279            self.meta_parameters[index] =
280                self.meta_parameters[index] - rate * self.meta_gradient_buffer[index];
281        }
282        Ok(())
283    }
284}
285
286impl<T: Float + Debug + Send + Sync + 'static> TaskDetector<T> {
287    /// Number of rounds currently buffered.
288    pub fn buffered_rounds(&self) -> usize {
289        self.gradient_buffer.len()
290    }
291
292    /// Detect whether this round's updates come from a different task.
293    ///
294    /// Returns `false` (and buffers the round) until there is history to compare
295    /// against. The comparison is the larger of
296    ///
297    /// * the relative L2 shift `||mean_now - mean_history|| / max(||mean_history||, eps)`, and
298    /// * the cosine dissimilarity `(1 - cos(mean_now, mean_history)) / 2`,
299    ///
300    /// so both a magnitude jump and a direction reversal are detected.
301    pub fn detect_task_change(&mut self, updates: &[Array1<T>], round: usize) -> Result<bool> {
302        match self.detection_method {
303            TaskDetectionMethod::GradientBased => {}
304            other => {
305                return Err(OptimError::UnsupportedOperation(format!(
306                    "TaskDetectionMethod::{other:?} is not implemented; only GradientBased \
307                     detection exists, and reporting its verdict under another name would \
308                     misdescribe what was measured"
309                )))
310            }
311        }
312        if updates.is_empty() {
313            return Err(OptimError::InvalidParameter(
314                "no client updates were supplied, so no task change can be detected".to_string(),
315            ));
316        }
317        let dimension = updates[0].len();
318        if dimension == 0 {
319            return Err(OptimError::InvalidParameter(
320                "the client updates are zero-dimensional".to_string(),
321            ));
322        }
323        for (index, update) in updates.iter().enumerate() {
324            if update.len() != dimension {
325                return Err(OptimError::DimensionMismatch(format!(
326                    "update {index} has {} elements, expected {dimension}",
327                    update.len()
328                )));
329            }
330        }
331
332        // Mean update for this round.
333        let mut mean_now = vec![0.0f64; dimension];
334        for update in updates {
335            let values = as_finite_f64("client update", update)?;
336            for (slot, value) in mean_now.iter_mut().zip(values.iter()) {
337                *slot += value;
338            }
339        }
340        let count = updates.len() as f64;
341        for slot in mean_now.iter_mut() {
342            *slot /= count;
343        }
344
345        // Running mean of the buffered history.
346        let history: Vec<Vec<f64>> = self
347            .gradient_buffer
348            .iter()
349            .map(|entry| as_finite_f64("buffered update", entry))
350            .collect::<Result<Vec<Vec<f64>>>>()?;
351        let comparable: Vec<&Vec<f64>> = history
352            .iter()
353            .filter(|entry| entry.len() == dimension)
354            .collect();
355
356        let mut detected = false;
357        if !comparable.is_empty() {
358            let mut mean_history = vec![0.0f64; dimension];
359            for entry in &comparable {
360                for (slot, value) in mean_history.iter_mut().zip(entry.iter()) {
361                    *slot += value;
362                }
363            }
364            let history_count = comparable.len() as f64;
365            for slot in mean_history.iter_mut() {
366                *slot /= history_count;
367            }
368
369            let history_norm: f64 = mean_history
370                .iter()
371                .map(|value| value * value)
372                .sum::<f64>()
373                .sqrt();
374            let shift_norm: f64 = mean_now
375                .iter()
376                .zip(mean_history.iter())
377                .map(|(now, past)| (now - past) * (now - past))
378                .sum::<f64>()
379                .sqrt();
380            let relative_shift = shift_norm / history_norm.max(f64::EPSILON);
381            let dissimilarity = cosine_similarity(&mean_now, &mean_history)
382                .map(|similarity| (1.0 - similarity) / 2.0)
383                .unwrap_or(0.0);
384            let magnitude = relative_shift.max(dissimilarity);
385
386            if magnitude > self.detection_threshold {
387                detected = true;
388                // Confidence rises with how far past the threshold the shift is,
389                // saturating at 1; it is a monotone transform of the measured
390                // magnitude, not an invented constant.
391                let excess = (magnitude - self.detection_threshold)
392                    / self.detection_threshold.max(f64::EPSILON);
393                self.change_points.push(ChangePoint {
394                    round,
395                    confidence: (excess / (1.0 + excess)).clamp(0.0, 1.0),
396                    change_magnitude: magnitude,
397                });
398            }
399        }
400
401        // Buffer this round's mean for the next comparison.
402        let mut buffered = Array1::zeros(dimension);
403        for (index, value) in mean_now.iter().enumerate() {
404            buffered[index] = T::from(*value).unwrap_or_else(T::zero);
405        }
406        self.gradient_buffer.push_back(buffered);
407        while self.gradient_buffer.len() > MAX_GRADIENT_BUFFER {
408            let _ = self.gradient_buffer.pop_front();
409        }
410
411        // A detected change starts a new task, so the old history is no longer
412        // the right baseline.
413        if detected {
414            self.gradient_buffer.clear();
415            let mut restart = Array1::zeros(dimension);
416            for (index, value) in mean_now.iter().enumerate() {
417                restart[index] = T::from(*value).unwrap_or_else(T::zero);
418            }
419            self.gradient_buffer.push_back(restart);
420        }
421
422        Ok(detected)
423    }
424}
425
426#[cfg(test)]
427mod tests {
428    use super::*;
429
430    fn array(values: &[f64]) -> Array1<f64> {
431        Array1::from(values.to_vec())
432    }
433
434    /// `(client gradients, support gradients, query gradients)`.
435    type GradientMaps = (
436        HashMap<String, Array1<f64>>,
437        HashMap<String, Array1<f64>>,
438        HashMap<String, Array1<f64>>,
439    );
440
441    /// `(name, reported gradient, support gradient, query gradient)`.
442    type ClientEntry<'a> = (&'a str, [f64; 3], [f64; 3], [f64; 3]);
443
444    fn maps(entries: &[ClientEntry<'_>]) -> GradientMaps {
445        let mut gradients = HashMap::new();
446        let mut support = HashMap::new();
447        let mut query = HashMap::new();
448        for (name, gradient, support_gradient, query_gradient) in entries {
449            gradients.insert((*name).to_string(), array(gradient));
450            support.insert((*name).to_string(), array(support_gradient));
451            query.insert((*name).to_string(), array(query_gradient));
452        }
453        (gradients, support, query)
454    }
455
456    #[test]
457    fn a_meta_learner_sized_for_a_model_allocates_real_buffers() {
458        // Regression for F109: `new` discarded its argument and allocated
459        // `Array1::default(0)`, so a caller sizing for a real model got nothing.
460        let learner = FederatedMetaLearner::<f64>::new(1_000);
461        assert_eq!(learner.parameter_size(), 1_000);
462        assert_eq!(learner.meta_parameters().len(), 1_000);
463        assert_eq!(learner.meta_gradient_buffer().len(), 1_000);
464    }
465
466    #[test]
467    fn an_unsized_meta_learner_reports_it_instead_of_returning_an_empty_array() {
468        let mut learner = FederatedMetaLearner::<f64>::new(0);
469        let (gradients, support, query) =
470            maps(&[("a", [1.0, 2.0, 3.0], [1.0, 1.0, 1.0], [0.5, 0.5, 0.5])]);
471        let message = match learner.compute_client_meta_gradients(&gradients, &support, &query) {
472            Err(err) => err.to_string(),
473            Ok(gradient) => panic!("returned a {}-element gradient", gradient.len()),
474        };
475        assert!(message.contains("parameter_size = 0"), "got: {message}");
476    }
477
478    #[test]
479    fn the_meta_gradient_is_the_mean_of_the_query_gradients() {
480        let mut learner = FederatedMetaLearner::<f64>::new(3);
481        let (gradients, support, query) = maps(&[
482            ("a", [1.0, 0.0, 0.0], [1.0, 0.0, 0.0], [2.0, 0.0, 0.0]),
483            ("b", [0.0, 1.0, 0.0], [0.0, 1.0, 0.0], [0.0, 4.0, 0.0]),
484        ]);
485        let meta_gradient =
486            match learner.compute_client_meta_gradients(&gradients, &support, &query) {
487                Ok(gradient) => gradient,
488                Err(err) => panic!("meta-gradient failed: {err}"),
489            };
490        assert_eq!(meta_gradient.len(), 3);
491        assert!((meta_gradient[0] - 1.0).abs() < 1e-12, "{meta_gradient:?}");
492        assert!((meta_gradient[1] - 2.0).abs() < 1e-12, "{meta_gradient:?}");
493        assert!((meta_gradient[2] - 0.0).abs() < 1e-12, "{meta_gradient:?}");
494        assert_eq!(learner.meta_gradient_buffer(), &meta_gradient);
495    }
496
497    #[test]
498    fn the_inner_step_is_recorded_per_client() {
499        let mut learner = FederatedMetaLearner::<f64>::new(3);
500        let (gradients, support, query) =
501            maps(&[("a", [1.0, 1.0, 1.0], [10.0, 20.0, 30.0], [1.0, 1.0, 1.0])]);
502        let ok = learner.compute_client_meta_gradients_with_rate(&gradients, &support, &query, 0.1);
503        assert!(ok.is_ok(), "meta-gradient failed");
504        let adaptation = match learner.client_adaptation("a") {
505            Some(adaptation) => adaptation,
506            None => panic!("the per-client adaptation must be recorded"),
507        };
508        // theta = 0 initially, so theta_a = -0.1 * [10, 20, 30].
509        assert!((adaptation[0] + 1.0).abs() < 1e-12, "{adaptation:?}");
510        assert!((adaptation[1] + 2.0).abs() < 1e-12, "{adaptation:?}");
511        assert!((adaptation[2] + 3.0).abs() < 1e-12, "{adaptation:?}");
512
513        let distribution = match learner.task_distribution("a") {
514            Some(distribution) => distribution,
515            None => panic!("the task distribution must be recorded"),
516        };
517        assert_eq!(distribution.adaptation_steps, 1);
518        // support = [10,20,30], query = [1,1,1]; the cosine is positive.
519        assert!(distribution.task_similarity > 0.0);
520        assert!(distribution.task_similarity <= 1.0);
521    }
522
523    #[test]
524    fn the_meta_gradient_is_reproducible_regardless_of_map_insertion_order() {
525        // `HashMap` iteration order is randomised per process, and float
526        // addition is not associative, so the sum must be over a sorted order.
527        let mut left = FederatedMetaLearner::<f64>::new(3);
528        let (gradients, support, query) = maps(&[
529            ("a", [1.0, 0.0, 0.0], [1.0, 0.0, 0.0], [0.1, 0.2, 0.3]),
530            ("b", [0.0, 1.0, 0.0], [0.0, 1.0, 0.0], [0.4, 0.5, 0.6]),
531            ("c", [0.0, 0.0, 1.0], [0.0, 0.0, 1.0], [0.7, 0.8, 0.9]),
532        ]);
533        let first = match left.compute_client_meta_gradients(&gradients, &support, &query) {
534            Ok(gradient) => gradient,
535            Err(err) => panic!("meta-gradient failed: {err}"),
536        };
537
538        let mut right = FederatedMetaLearner::<f64>::new(3);
539        let mut reordered_gradients = HashMap::new();
540        for name in ["c", "a", "b"] {
541            if let Some(value) = gradients.get(name) {
542                reordered_gradients.insert(name.to_string(), value.clone());
543            }
544        }
545        let second =
546            match right.compute_client_meta_gradients(&reordered_gradients, &support, &query) {
547                Ok(gradient) => gradient,
548                Err(err) => panic!("meta-gradient failed: {err}"),
549            };
550        assert_eq!(first, second);
551    }
552
553    #[test]
554    fn a_mismatched_or_missing_client_gradient_is_refused() {
555        let mut learner = FederatedMetaLearner::<f64>::new(3);
556        let (gradients, support, query) =
557            maps(&[("a", [1.0, 2.0, 3.0], [1.0, 1.0, 1.0], [0.5, 0.5, 0.5])]);
558
559        // Wrong length.
560        let mut wrong = gradients.clone();
561        wrong.insert("a".to_string(), array(&[1.0, 2.0]));
562        assert!(learner
563            .compute_client_meta_gradients(&wrong, &support, &query)
564            .is_err());
565
566        // Missing support gradient.
567        assert!(learner
568            .compute_client_meta_gradients(&gradients, &HashMap::new(), &query)
569            .is_err());
570        // Missing query gradient.
571        assert!(learner
572            .compute_client_meta_gradients(&gradients, &support, &HashMap::new())
573            .is_err());
574        // No clients at all.
575        assert!(learner
576            .compute_client_meta_gradients(&HashMap::new(), &support, &query)
577            .is_err());
578        // Non-finite input.
579        let mut poisoned = support.clone();
580        poisoned.insert("a".to_string(), array(&[f64::NAN, 1.0, 1.0]));
581        assert!(learner
582            .compute_client_meta_gradients(&gradients, &poisoned, &query)
583            .is_err());
584        // Bad inner learning rate.
585        assert!(learner
586            .compute_client_meta_gradients_with_rate(&gradients, &support, &query, 0.0)
587            .is_err());
588    }
589
590    #[test]
591    fn applying_the_meta_gradient_moves_the_meta_parameters() {
592        let mut learner = FederatedMetaLearner::<f64>::new(2);
593        let mut gradients = HashMap::new();
594        gradients.insert("a".to_string(), array(&[1.0, 1.0]));
595        let mut support = HashMap::new();
596        support.insert("a".to_string(), array(&[1.0, 1.0]));
597        let mut query = HashMap::new();
598        query.insert("a".to_string(), array(&[2.0, -4.0]));
599        let ok = learner.compute_client_meta_gradients(&gradients, &support, &query);
600        assert!(ok.is_ok());
601
602        let ok = learner.apply_meta_gradient(0.5);
603        assert!(ok.is_ok(), "apply failed");
604        assert!((learner.meta_parameters()[0] + 1.0).abs() < 1e-12);
605        assert!((learner.meta_parameters()[1] - 2.0).abs() < 1e-12);
606        assert!(learner.apply_meta_gradient(0.0).is_err());
607        assert!(learner.apply_meta_gradient(f64::NAN).is_err());
608    }
609
610    #[test]
611    fn a_stable_gradient_stream_reports_no_task_change() {
612        // Regression: the placeholder returned `Ok(false)` for everything, so a
613        // test that only checks "no change on stable input" would have passed
614        // against it. The next test is the one that could not.
615        let mut detector = TaskDetector::<f64>::new();
616        for round in 0..8usize {
617            let updates = vec![array(&[1.0, 0.5, -0.25]), array(&[1.02, 0.48, -0.26])];
618            match detector.detect_task_change(&updates, round) {
619                Ok(false) => {}
620                Ok(true) => panic!("a stable stream must not flag a change at round {round}"),
621                Err(err) => panic!("detection failed: {err}"),
622            }
623        }
624        assert!(detector.change_points().is_empty());
625        assert!(detector.buffered_rounds() > 0);
626    }
627
628    #[test]
629    fn a_direction_reversal_is_detected() {
630        let mut detector = TaskDetector::<f64>::new();
631        for round in 0..4usize {
632            let updates = vec![array(&[1.0, 1.0, 1.0])];
633            match detector.detect_task_change(&updates, round) {
634                Ok(false) => {}
635                Ok(true) => panic!("the warm-up rounds must not flag a change"),
636                Err(err) => panic!("detection failed: {err}"),
637            }
638        }
639        // Same magnitude, opposite direction.
640        let flipped = vec![array(&[-1.0, -1.0, -1.0])];
641        match detector.detect_task_change(&flipped, 4) {
642            Ok(true) => {}
643            Ok(false) => panic!("a sign flip must be detected"),
644            Err(err) => panic!("detection failed: {err}"),
645        }
646        let change = match detector.change_points().first() {
647            Some(change) => change,
648            None => panic!("the change point must be recorded"),
649        };
650        assert_eq!(change.round, 4);
651        assert!(change.change_magnitude > detector.detection_threshold());
652        assert!((0.0..=1.0).contains(&change.confidence));
653    }
654
655    #[test]
656    fn a_magnitude_jump_is_detected() {
657        let mut detector = TaskDetector::<f64>::new();
658        for round in 0..4usize {
659            let updates = vec![array(&[0.01, 0.01, 0.01])];
660            let ok = detector.detect_task_change(&updates, round);
661            assert!(matches!(ok, Ok(false)), "{ok:?}");
662        }
663        let jump = vec![array(&[10.0, 10.0, 10.0])];
664        match detector.detect_task_change(&jump, 4) {
665            Ok(true) => {}
666            Ok(false) => panic!("a 1000x magnitude jump must be detected"),
667            Err(err) => panic!("detection failed: {err}"),
668        }
669    }
670
671    #[test]
672    fn the_threshold_governs_sensitivity() {
673        let mut sensitive = TaskDetector::<f64>::new();
674        let ok = sensitive.set_detection_threshold(0.001);
675        assert!(ok.is_ok());
676        let mut tolerant = TaskDetector::<f64>::new();
677        let ok = tolerant.set_detection_threshold(5.0);
678        assert!(ok.is_ok());
679
680        for round in 0..3usize {
681            let updates = vec![array(&[1.0, 1.0])];
682            let _ = sensitive.detect_task_change(&updates, round);
683            let _ = tolerant.detect_task_change(&updates, round);
684        }
685        let shifted = vec![array(&[1.2, 1.2])];
686        assert!(
687            matches!(sensitive.detect_task_change(&shifted, 3), Ok(true)),
688            "a 0.001 threshold must flag a 20% shift"
689        );
690        assert!(
691            matches!(tolerant.detect_task_change(&shifted, 3), Ok(false)),
692            "a 5.0 threshold must not flag a 20% shift"
693        );
694        assert!(tolerant.set_detection_threshold(0.0).is_err());
695        assert!(tolerant.set_detection_threshold(f64::NAN).is_err());
696    }
697
698    #[test]
699    fn detection_restarts_its_baseline_after_a_change() {
700        let mut detector = TaskDetector::<f64>::new();
701        for round in 0..3usize {
702            let _ = detector.detect_task_change(&[array(&[1.0, 1.0])], round);
703        }
704        assert!(matches!(
705            detector.detect_task_change(&[array(&[-1.0, -1.0])], 3),
706            Ok(true)
707        ));
708        assert_eq!(
709            detector.buffered_rounds(),
710            1,
711            "the baseline must restart from the new task"
712        );
713        // The new regime is now normal, so it must not keep firing.
714        assert!(matches!(
715            detector.detect_task_change(&[array(&[-1.0, -1.0])], 4),
716            Ok(false)
717        ));
718        assert_eq!(detector.change_points().len(), 1);
719    }
720
721    #[test]
722    fn degenerate_detection_inputs_are_refused() {
723        let mut detector = TaskDetector::<f64>::new();
724        assert!(detector.detect_task_change(&[], 0).is_err());
725        assert!(detector
726            .detect_task_change(&[Array1::<f64>::zeros(0)], 0)
727            .is_err());
728        assert!(detector
729            .detect_task_change(&[array(&[1.0, 2.0]), array(&[1.0])], 0)
730            .is_err());
731        assert!(detector
732            .detect_task_change(&[array(&[f64::INFINITY, 1.0])], 0)
733            .is_err());
734    }
735
736    #[test]
737    fn unimplemented_detection_methods_are_refused() {
738        for method in [
739            TaskDetectionMethod::LossBased,
740            TaskDetectionMethod::StatisticalTest,
741            TaskDetectionMethod::ChangePointDetection,
742            TaskDetectionMethod::EnsembleMethods,
743        ] {
744            let mut detector = TaskDetector::<f64>::new();
745            detector.detection_method = method;
746            let outcome = detector.detect_task_change(&[array(&[1.0, 1.0])], 0);
747            assert!(
748                outcome.is_err(),
749                "{method:?} must not report a gradient-based verdict under another name"
750            );
751        }
752    }
753}