1use 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
49pub const DEFAULT_INNER_LEARNING_RATE: f64 = 0.01;
51
52const MAX_GRADIENT_BUFFER: usize = 100;
54
55fn 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
79fn 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 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 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 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 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 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 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 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 pub fn buffered_rounds(&self) -> usize {
289 self.gradient_buffer.len()
290 }
291
292 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 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 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 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 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 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 type GradientMaps = (
436 HashMap<String, Array1<f64>>,
437 HashMap<String, Array1<f64>>,
438 HashMap<String, Array1<f64>>,
439 );
440
441 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 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 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 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 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 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 assert!(learner
568 .compute_client_meta_gradients(&gradients, &HashMap::new(), &query)
569 .is_err());
570 assert!(learner
572 .compute_client_meta_gradients(&gradients, &support, &HashMap::new())
573 .is_err());
574 assert!(learner
576 .compute_client_meta_gradients(&HashMap::new(), &support, &query)
577 .is_err());
578 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 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 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 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 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}