1use crate::data::TimeSeries;
2use crate::error::EvaluatorError;
3use crate::evaluator::*;
4use crate::feature::Feature;
5use crate::float_trait::Float;
6
7use std::marker::PhantomData;
8
9macro_const! {
10 const DOC: &str = r#"
11Bulk feature extractor
12
13- Depends on: as reuired by feature evaluators
14- Minimum number of observations: as required by feature evaluators
15- Number of features: total for all feature evaluators
16"#;
17}
18
19#[doc = DOC!()]
20#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
21#[serde(
22 into = "FeatureExtractorParameters<F>",
23 from = "FeatureExtractorParameters<F>",
24 bound = "T: Float, F: FeatureEvaluator<T>"
25)]
26pub struct FeatureExtractor<T, F> {
27 features: Vec<F>,
28 info: Box<EvaluatorInfo>,
29 phantom: PhantomData<T>,
30}
31
32impl<T, F> FeatureExtractor<T, F>
33where
34 T: Float,
35 F: FeatureEvaluator<T>,
36{
37 pub fn new(features: Vec<F>) -> Self {
38 let info = EvaluatorInfo {
39 size: features.iter().map(|x| x.size_hint()).sum(),
40 min_ts_length: features
41 .iter()
42 .map(|x| x.min_ts_length())
43 .max()
44 .unwrap_or(0),
45 t_required: features.iter().any(|x| x.is_t_required()),
46 m_required: features.iter().any(|x| x.is_m_required()),
47 w_required: features.iter().any(|x| x.is_w_required()),
48 sorting_required: features.iter().any(|x| x.is_sorting_required()),
49 variability_required: features.iter().any(|x| x.is_variability_required()),
50 }
51 .into();
52 Self {
53 info,
54 features,
55 phantom: PhantomData,
56 }
57 }
58
59 pub fn get_features(&self) -> &Vec<F> {
60 &self.features
61 }
62
63 pub fn into_vec(self) -> Vec<F> {
64 self.features
65 }
66
67 pub fn add_feature(&mut self, feature: F) {
68 self.info.size += feature.size_hint();
69 self.info.min_ts_length = self.info.min_ts_length.max(feature.min_ts_length());
70 self.info.t_required |= feature.is_t_required();
71 self.info.m_required |= feature.is_m_required();
72 self.info.w_required |= feature.is_w_required();
73 self.info.sorting_required |= feature.is_sorting_required();
74 self.info.variability_required |= feature.is_variability_required();
75 self.features.push(feature);
76 }
77}
78
79impl<T> FeatureExtractor<T, Feature<T>>
80where
81 T: Float,
82{
83 pub fn from_features(features: Vec<Feature<T>>) -> Self {
85 Self::new(features)
86 }
87}
88
89impl<T, F> FeatureExtractor<T, F> {
90 pub const fn doc() -> &'static str {
91 DOC
92 }
93}
94
95impl<T, F> EvaluatorInfoTrait for FeatureExtractor<T, F>
96where
97 T: Float,
98 F: FeatureEvaluator<T>,
99{
100 fn get_info(&self) -> &EvaluatorInfo {
101 &self.info
102 }
103}
104
105impl<T, F> FeatureNamesDescriptionsTrait for FeatureExtractor<T, F>
106where
107 T: Float,
108 F: FeatureEvaluator<T>,
109{
110 fn get_names(&self) -> Vec<&str> {
112 self.features.iter().flat_map(|x| x.get_names()).collect()
113 }
114
115 fn get_descriptions(&self) -> Vec<&str> {
117 self.features
118 .iter()
119 .flat_map(|x| x.get_descriptions())
120 .collect()
121 }
122}
123
124impl<T, F> FeatureEvaluator<T> for FeatureExtractor<T, F>
125where
126 T: Float,
127 F: FeatureEvaluator<T>,
128{
129 fn eval_no_ts_check(&self, ts: &mut TimeSeries<T>) -> Result<Vec<T>, EvaluatorError> {
130 let mut vec = Vec::with_capacity(self.size_hint());
131 for x in &self.features {
132 vec.extend(x.eval_no_ts_check(ts)?);
133 }
134 Ok(vec)
135 }
136
137 fn eval_or_fill(&self, ts: &mut TimeSeries<T>, fill_value: T) -> Vec<T> {
138 self.features
139 .iter()
140 .flat_map(|x| x.eval_or_fill(ts, fill_value))
141 .collect()
142 }
143}
144
145#[cfg(test)]
146impl<T, F> Default for FeatureExtractor<T, F>
147where
148 T: Float,
149 F: FeatureEvaluator<T>,
150{
151 fn default() -> Self {
152 Self::new(vec![])
153 }
154}
155
156#[derive(Serialize, Deserialize, JsonSchema)]
157#[serde(rename = "FeatureExtractor")]
158struct FeatureExtractorParameters<F> {
159 features: Vec<F>,
160}
161
162impl<T, F> From<FeatureExtractor<T, F>> for FeatureExtractorParameters<F> {
163 fn from(f: FeatureExtractor<T, F>) -> Self {
164 Self {
165 features: f.features,
166 }
167 }
168}
169
170impl<T, F> From<FeatureExtractorParameters<F>> for FeatureExtractor<T, F>
171where
172 T: Float,
173 F: FeatureEvaluator<T>,
174{
175 fn from(p: FeatureExtractorParameters<F>) -> Self {
176 Self::new(p.features)
177 }
178}
179
180impl<T, F> JsonSchema for FeatureExtractor<T, F>
181where
182 F: JsonSchema,
183{
184 json_schema!(FeatureExtractorParameters<F>, true);
185}
186
187#[cfg(test)]
188mod tests {
189 use super::*;
190 use crate::Feature;
191 use crate::tests::*;
192
193 use approx::assert_relative_eq;
194 use serde_test::{Token, assert_ser_tokens};
195
196 serialization_name_test!(FeatureExtractor<f64, Feature<f64>>);
197
198 serde_json_test!(
199 feature_extractor_ser_json_de,
200 FeatureExtractor<f64, Feature<f64>>,
201 FeatureExtractor::new(vec![crate::Amplitude{}.into(), crate::BeyondNStd::new(2.0).into()]),
202 );
203
204 check_doc_static_method!(feature_extractor_doc_static_method, FeatureExtractor<f64, Feature<f64>>);
205
206 #[test]
207 fn serialization_empty() {
208 let fe: FeatureExtractor<f64, Feature<_>> = FeatureExtractor::new(vec![]);
209 assert_ser_tokens(
210 &fe,
211 &[
212 Token::Struct {
214 len: 1,
215 name: "FeatureExtractor",
216 },
217 Token::String("features"),
219 Token::Seq { len: Some(0) },
220 Token::SeqEnd,
221 Token::StructEnd,
223 ],
224 )
225 }
226
227 #[test]
229 fn multi_feature_eval_values() {
230 let t = [0.0_f64, 1.0, 2.0, 3.0, 4.0];
231 let m = [1.0_f64, 2.0, 3.0, 4.0, 5.0];
232 let w = [1.0_f64; 5];
233 let mut ts = TimeSeries::new(&t[..], &m[..], &w[..]);
234
235 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
236 crate::Amplitude::new().into(),
237 crate::Mean::new().into(),
238 ]);
239
240 let values = fe.eval(&mut ts).unwrap();
241 assert_eq!(values.len(), 2, "should produce one value per feature");
242 assert_relative_eq!(values[0], 2.0, epsilon = 1e-10);
244 assert_relative_eq!(values[1], 3.0, epsilon = 1e-10);
246 }
247
248 #[test]
250 fn multi_feature_names_and_descriptions_aggregated() {
251 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
252 crate::Amplitude::new().into(),
253 crate::Mean::new().into(),
254 crate::StandardDeviation::new().into(),
255 ]);
256
257 let names = fe.get_names();
258 let descs = fe.get_descriptions();
259
260 assert_eq!(names.len(), 3);
261 assert_eq!(descs.len(), 3);
262 assert_eq!(fe.size_hint(), 3);
263 assert_eq!(names[0], "amplitude");
265 assert_eq!(names[1], "mean");
266 assert_eq!(names[2], "standard_deviation");
267 assert!(descs.iter().all(|d| !d.is_empty()));
269 }
270
271 #[test]
273 fn info_aggregated_correctly() {
274 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
277 crate::Amplitude::new().into(),
278 crate::LinearTrend::new().into(),
279 ]);
280
281 assert!(
282 fe.is_t_required(),
283 "t_required should be true when any feature requires it"
284 );
285 assert!(
286 fe.is_sorting_required(),
287 "sorting_required should be true when any feature requires it"
288 );
289 assert_eq!(
290 fe.min_ts_length(),
291 3,
292 "min_ts_length should be the maximum across features"
293 );
294 assert_eq!(
295 fe.size_hint(),
296 1 + 3,
297 "size should be the sum across features"
298 );
299 }
300
301 #[test]
303 fn add_feature_updates_info_correctly() {
304 let mut fe: FeatureExtractor<f64, Feature<f64>> =
305 FeatureExtractor::new(vec![crate::Amplitude::new().into()]);
306
307 assert_eq!(fe.size_hint(), 1);
308 assert!(!fe.is_t_required());
309 assert!(!fe.is_sorting_required());
310 assert_eq!(fe.min_ts_length(), 1);
311
312 fe.add_feature(crate::LinearTrend::new().into());
313
314 assert_eq!(fe.size_hint(), 4);
315 assert!(fe.is_t_required());
316 assert!(fe.is_sorting_required());
317 assert_eq!(fe.min_ts_length(), 3);
318 }
319
320 #[test]
322 fn eval_returns_error_on_short_ts() {
323 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
325 crate::Amplitude::new().into(),
326 crate::LinearTrend::new().into(),
327 ]);
328
329 let t = [0.0_f64, 1.0];
330 let m = [1.0_f64, 2.0];
331 let w = [1.0_f64, 1.0];
332 let mut ts = TimeSeries::new(&t[..], &m[..], &w[..]);
333
334 let result = fe.eval(&mut ts);
335 assert!(
336 matches!(
337 result,
338 Err(EvaluatorError::ShortTimeSeries {
339 actual: 2,
340 minimum: 3
341 })
342 ),
343 "expected ShortTimeSeries error, got: {:?}",
344 result
345 );
346 }
347
348 #[test]
352 fn eval_or_fill_fills_only_failing_feature() {
353 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
355 crate::Amplitude::new().into(), crate::LinearTrend::new().into(), ]);
358
359 let t = [0.0_f64, 1.0];
360 let m = [1.0_f64, 3.0];
361 let w = [1.0_f64, 1.0];
362 let mut ts = TimeSeries::new(&t[..], &m[..], &w[..]);
363
364 let values = fe.eval_or_fill(&mut ts, f64::NAN);
365 assert_eq!(values.len(), 4, "should always return size_hint() values");
366 assert_relative_eq!(values[0], 1.0, epsilon = 1e-10);
368 assert!(
370 values[1..].iter().all(|v| v.is_nan()),
371 "failed feature outputs should be fill value"
372 );
373 }
374
375 #[test]
377 fn eval_or_fill_fills_all_on_single_failing_feature() {
378 let fe: FeatureExtractor<f64, Feature<f64>> =
380 FeatureExtractor::new(vec![crate::OtsuSplit::new().into()]);
381
382 let t = [0.0_f64, 1.0, 2.0, 3.0];
383 let m = [3.0_f64; 4];
384 let w = [1.0_f64; 4];
385 let mut ts = TimeSeries::new(&t[..], &m[..], &w[..]);
386
387 let values = fe.eval_or_fill(&mut ts, -999.0);
388 assert_eq!(values.len(), fe.size_hint());
389 assert!(
390 values.iter().all(|&v| v == -999.0),
391 "all outputs should be fill value"
392 );
393 }
394
395 #[test]
397 fn eval_or_fill_returns_values_on_valid_ts() {
398 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
399 crate::Amplitude::new().into(),
400 crate::Mean::new().into(),
401 ]);
402
403 let t = [0.0_f64, 1.0, 2.0, 3.0, 4.0];
404 let m = [1.0_f64, 2.0, 3.0, 4.0, 5.0];
405 let w = [1.0_f64; 5];
406 let mut ts = TimeSeries::new(&t[..], &m[..], &w[..]);
407
408 let values = fe.eval_or_fill(&mut ts, f64::NAN);
409 assert_eq!(values.len(), 2);
410 assert!(
411 values.iter().all(|v| v.is_finite()),
412 "values should be finite"
413 );
414 }
415
416 #[test]
418 fn eval_result_length_consistent_with_size_hint() {
419 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
420 crate::Amplitude::new().into(),
421 crate::LinearTrend::new().into(),
422 crate::Mean::new().into(),
423 ]);
424
425 let t = [0.0_f64, 1.0, 2.0, 3.0, 4.0];
426 let m = [1.0_f64, 3.0, 2.0, 5.0, 4.0];
427 let w = [1.0_f64; 5];
428 let mut ts = TimeSeries::new(&t[..], &m[..], &w[..]);
429
430 let values = fe.eval(&mut ts).unwrap();
431 assert_eq!(values.len(), fe.size_hint());
432 assert_eq!(values.len(), fe.get_names().len());
433 assert_eq!(values.len(), fe.get_descriptions().len());
434 }
435
436 #[test]
438 fn eval_returns_flat_ts_error_for_constant_magnitude() {
439 let fe: FeatureExtractor<f64, Feature<f64>> =
441 FeatureExtractor::new(vec![crate::OtsuSplit::new().into()]);
442
443 let t = [0.0_f64, 1.0, 2.0, 3.0];
444 let m = [3.0_f64; 4];
445 let w = [1.0_f64; 4];
446 let mut ts = TimeSeries::new(&t[..], &m[..], &w[..]);
447
448 assert!(
449 matches!(fe.eval(&mut ts), Err(EvaluatorError::FlatTimeSeries)),
450 "expected FlatTimeSeries error for constant magnitude input"
451 );
452 }
453
454 #[test]
456 fn full_pipeline_on_realistic_data() {
457 let mut rng = StdRng::seed_from_u64(42);
458 let n = 50;
459 let t: Vec<f64> = sorted(&randvec::<f64>(&mut rng, n))
460 .into_iter()
461 .enumerate()
462 .map(|(i, _)| i as f64)
463 .collect();
464 let m = randvec::<f64>(&mut rng, n);
465 let w = positive_randvec::<f64>(&mut rng, n);
466
467 let fe: FeatureExtractor<f64, Feature<f64>> = FeatureExtractor::new(vec![
468 crate::Amplitude::new().into(),
469 crate::Mean::new().into(),
470 crate::StandardDeviation::new().into(),
471 crate::LinearTrend::new().into(),
472 crate::MedianAbsoluteDeviation::new().into(),
473 ]);
474
475 let expected_size = fe.size_hint();
476 let mut ts = TimeSeries::new(&t, &m, &w);
477 let values = fe.eval(&mut ts).unwrap();
478
479 assert_eq!(values.len(), expected_size);
480 assert_eq!(values.len(), fe.get_names().len());
481 assert!(
482 values.iter().all(|v| v.is_finite()),
483 "all pipeline values should be finite"
484 );
485 }
486}