1use crate::error::{RillError, checked_increment, ensure_finite, validate_features};
7use crate::traits::Transformer;
8
9#[derive(Debug, Clone)]
11#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
12pub struct StandardScalerConfig {
13 pub with_mean: bool,
15 pub with_std: bool,
17 pub epsilon: f64,
20}
21
22impl Default for StandardScalerConfig {
23 fn default() -> Self {
24 Self {
25 with_mean: true,
26 with_std: true,
27 epsilon: 1e-12,
28 }
29 }
30}
31
32#[derive(Debug, Clone)]
44#[cfg_attr(feature = "serde", derive(serde::Serialize))]
45pub struct StandardScaler {
46 feature_count: usize,
47 config: StandardScalerConfig,
48 counts: Vec<u64>,
49 means: Vec<f64>,
50 m2s: Vec<f64>,
51}
52
53impl StandardScaler {
54 pub fn new(feature_count: usize) -> Result<Self, RillError> {
56 Self::with_config(feature_count, StandardScalerConfig::default())
57 }
58
59 pub fn with_config(
61 feature_count: usize,
62 config: StandardScalerConfig,
63 ) -> Result<Self, RillError> {
64 if feature_count == 0 {
65 return Err(RillError::EmptyFeatures);
66 }
67 ensure_finite("epsilon", config.epsilon)?;
68 if config.epsilon < 0.0 {
69 return Err(RillError::InvalidParameter {
70 name: "epsilon",
71 value: config.epsilon,
72 });
73 }
74 Ok(Self {
75 feature_count,
76 config,
77 counts: vec![0; feature_count],
78 means: vec![0.0; feature_count],
79 m2s: vec![0.0; feature_count],
80 })
81 }
82
83 pub fn means(&self) -> &[f64] {
85 &self.means
86 }
87
88 pub fn variances(&self) -> Vec<f64> {
90 self.m2s
91 .iter()
92 .zip(&self.counts)
93 .map(|(&m2, &n)| if n == 0 { 0.0 } else { m2 / n as f64 })
94 .collect()
95 }
96
97 pub fn std_devs(&self) -> Vec<f64> {
99 self.variances().iter().map(|v| v.sqrt()).collect()
100 }
101
102 pub fn scales(&self) -> Vec<f64> {
104 self.variances()
105 .iter()
106 .map(|&var| {
107 if var < self.config.epsilon {
108 1.0
109 } else {
110 var.sqrt()
111 }
112 })
113 .collect()
114 }
115
116 pub fn validate(&self) -> Result<(), RillError> {
123 if self.feature_count == 0 {
124 return Err(RillError::EmptyFeatures);
125 }
126 ensure_finite("epsilon", self.config.epsilon)?;
127 if self.config.epsilon < 0.0 {
128 return Err(RillError::InvalidParameter {
129 name: "epsilon",
130 value: self.config.epsilon,
131 });
132 }
133 if self.counts.len() != self.feature_count
134 || self.means.len() != self.feature_count
135 || self.m2s.len() != self.feature_count
136 {
137 return Err(RillError::InvalidState(
138 "standard scaler feature_count does not match state lengths".to_owned(),
139 ));
140 }
141 if self.means.iter().any(|value| !value.is_finite())
142 || self.m2s.iter().any(|value| !value.is_finite())
143 {
144 return Err(RillError::InvalidState(
145 "standard scaler state must contain only finite values".to_owned(),
146 ));
147 }
148 if self.m2s.iter().any(|value| *value < 0.0) {
154 return Err(RillError::InvalidState(
155 "standard scaler m2s must be non-negative".to_owned(),
156 ));
157 }
158 if self.counts.windows(2).any(|pair| pair[0] != pair[1]) {
159 return Err(RillError::InvalidState(
160 "standard scaler feature counts must stay synchronized".to_owned(),
161 ));
162 }
163 let n = self.counts.first().copied().unwrap_or(0);
166 if n == 0 {
172 for (i, &mean) in self.means.iter().enumerate() {
173 if mean != 0.0 {
174 return Err(RillError::InvalidState(format!(
175 "standard scaler means[{i}] must be 0 when count == 0, got {mean}"
176 )));
177 }
178 }
179 for (i, &m2) in self.m2s.iter().enumerate() {
180 if m2 != 0.0 {
181 return Err(RillError::InvalidState(format!(
182 "standard scaler m2s[{i}] must be 0 when count == 0, got {m2}"
183 )));
184 }
185 }
186 }
187 if n == 1 {
193 for (i, &m2) in self.m2s.iter().enumerate() {
194 if m2 != 0.0 {
195 return Err(RillError::InvalidState(format!(
196 "standard scaler m2s[{i}] must be 0 when count == 1, got {m2}"
197 )));
198 }
199 }
200 }
201 Ok(())
202 }
203
204 pub fn transform_into(&self, features: &[f64], output: &mut Vec<f64>) -> Result<(), RillError> {
220 validate_features(self.feature_count, features)?;
223
224 debug_assert!(
228 self.counts.len() == self.feature_count
229 && self.means.len() == self.feature_count
230 && self.m2s.len() == self.feature_count,
231 "standard scaler state lengths must match feature_count"
232 );
233 debug_assert!(
234 self.counts.windows(2).all(|pair| pair[0] == pair[1]),
235 "standard scaler feature counts must stay synchronized"
236 );
237 debug_assert!(
238 self.means.iter().all(|v| v.is_finite()) && self.m2s.iter().all(|v| v.is_finite()),
239 "standard scaler state must contain only finite values"
240 );
241
242 output.clear();
243 output.reserve(self.feature_count);
244
245 let iter = features
246 .iter()
247 .zip(&self.counts)
248 .zip(&self.means)
249 .zip(&self.m2s);
250 for (((&x, &n), &mean_storage), &m2) in iter {
251 let scale = if !self.config.with_std || n == 0 {
254 1.0
255 } else {
256 let variance = m2 / n as f64;
257 if variance < self.config.epsilon {
258 1.0
259 } else {
260 variance.sqrt()
261 }
262 };
263 let mean = if self.config.with_mean {
264 mean_storage
265 } else {
266 0.0
267 };
268 let transformed = (x - mean) / scale;
269 ensure_finite("transformed feature", transformed)?;
273 output.push(transformed);
274 }
275 Ok(())
276 }
277}
278
279impl Transformer for StandardScaler {
280 fn input_dim(&self) -> usize {
281 self.feature_count
282 }
283
284 fn output_dim(&self) -> usize {
285 self.feature_count
286 }
287
288 fn transform(&self, features: &[f64]) -> Result<Vec<f64>, RillError> {
289 let mut output = Vec::with_capacity(self.feature_count);
296 self.transform_into(features, &mut output)?;
297 Ok(output)
298 }
299
300 fn update(&mut self, features: &[f64]) -> Result<(), RillError> {
301 self.validate()?;
302 validate_features(self.feature_count, features)?;
303 let mut next_counts = self.counts.clone();
304 let mut next_means = self.means.clone();
305 let mut next_m2s = self.m2s.clone();
306 for (i, &x) in features.iter().enumerate() {
307 let count = checked_increment(self.counts[i], "standard scaler sample")?;
308 let delta = x - self.means[i];
309 ensure_finite("standard scaler delta", delta)?;
310 let mean = self.means[i] + delta / count as f64;
311 ensure_finite("standard scaler mean", mean)?;
312 let delta2 = x - mean;
313 ensure_finite("standard scaler delta", delta2)?;
314 let m2 = self.m2s[i] + delta * delta2;
315 ensure_finite("standard scaler M2", m2)?;
316 next_counts[i] = count;
317 next_means[i] = mean;
318 next_m2s[i] = m2;
319 }
320 self.counts = next_counts;
321 self.means = next_means;
322 self.m2s = next_m2s;
323 Ok(())
324 }
325
326 fn samples_seen(&self) -> u64 {
327 self.counts.iter().copied().max().unwrap_or(0)
328 }
329
330 fn reset(&mut self) {
331 for c in &mut self.counts {
332 *c = 0;
333 }
334 for m in &mut self.means {
335 *m = 0.0;
336 }
337 for m2 in &mut self.m2s {
338 *m2 = 0.0;
339 }
340 }
341}
342
343#[cfg(feature = "serde")]
344#[derive(serde::Deserialize)]
345struct StandardScalerState {
346 feature_count: usize,
347 config: StandardScalerConfig,
348 counts: Vec<u64>,
349 means: Vec<f64>,
350 m2s: Vec<f64>,
351}
352
353#[cfg(feature = "serde")]
354impl<'de> serde::Deserialize<'de> for StandardScaler {
355 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
356 where
357 D: serde::Deserializer<'de>,
358 {
359 let state = StandardScalerState::deserialize(deserializer)?;
360 let scaler = Self {
361 feature_count: state.feature_count,
362 config: state.config,
363 counts: state.counts,
364 means: state.means,
365 m2s: state.m2s,
366 };
367 scaler.validate().map_err(serde::de::Error::custom)?;
368 Ok(scaler)
369 }
370}
371
372#[cfg(test)]
373mod tests {
374 use super::*;
375
376 #[test]
377 fn scaler_zero_state_returns_original() {
378 let s = StandardScaler::new(3).unwrap();
379 let out = s.transform(&[1.0, 2.0, 3.0]).unwrap();
380 assert!((out[0] - 1.0).abs() < 1e-12);
382 assert!((out[1] - 2.0).abs() < 1e-12);
383 assert!((out[2] - 3.0).abs() < 1e-12);
384 }
385
386 #[test]
387 fn scaler_standardizes_after_updates() {
388 let mut s = StandardScaler::new(2).unwrap();
389 s.update(&[1.0, 10.0]).unwrap();
392 s.update(&[3.0, 20.0]).unwrap();
393 let out = s.transform(&[3.0, 20.0]).unwrap();
394 assert!((out[0] - 1.0).abs() < 1e-9);
396 assert!((out[1] - 1.0).abs() < 1e-9);
397 }
398
399 #[test]
400 fn transform_does_not_update_state() {
401 let mut s = StandardScaler::new(1).unwrap();
402 s.update(&[10.0]).unwrap();
403 let mean_before = s.means()[0];
404 let _ = s.transform(&[5.0]).unwrap();
405 assert_eq!(s.means()[0], mean_before);
406 assert_eq!(s.counts[0], 1);
407 }
408
409 #[test]
410 fn update_rejects_overflow_without_mutating_state() {
411 let mut scaler = StandardScaler::new(1).unwrap();
412 scaler.update(&[f64::MAX]).unwrap();
413 let before = scaler.clone();
414 assert!(scaler.update(&[-f64::MAX]).is_err());
415 assert_eq!(scaler.counts, before.counts);
416 assert_eq!(scaler.means, before.means);
417 assert_eq!(scaler.m2s, before.m2s);
418 }
419
420 #[cfg(feature = "serde")]
421 #[test]
422 fn serde_rejects_malformed_state() {
423 let malformed = r#"{
424 "feature_count":2,
425 "config":{"with_mean":true,"with_std":true,"epsilon":1e-12},
426 "counts":[1],
427 "means":[0.0],
428 "m2s":[0.0]
429 }"#;
430 assert!(serde_json::from_str::<StandardScaler>(malformed).is_err());
431 }
432
433 #[cfg(feature = "serde")]
434 #[test]
435 fn serde_rejects_negative_m2() {
436 let malformed = r#"{
441 "feature_count":1,
442 "config":{"with_mean":true,"with_std":true,"epsilon":1e-12},
443 "counts":[2],
444 "means":[0.0],
445 "m2s":[-1.0]
446 }"#;
447 assert!(serde_json::from_str::<StandardScaler>(malformed).is_err());
448 }
449
450 #[cfg(feature = "serde")]
451 #[test]
452 fn scaler_serde_rejects_zero_count_nonzero_mean() {
453 let malformed = r#"{
456 "feature_count":1,
457 "config":{"with_mean":true,"with_std":true,"epsilon":1e-12},
458 "counts":[0],
459 "means":[10.0],
460 "m2s":[0.0]
461 }"#;
462 assert!(
463 serde_json::from_str::<StandardScaler>(malformed).is_err(),
464 "count=0 with non-zero mean must be rejected"
465 );
466 }
467
468 #[cfg(feature = "serde")]
469 #[test]
470 fn scaler_serde_rejects_zero_count_nonzero_m2() {
471 let malformed = r#"{
474 "feature_count":1,
475 "config":{"with_mean":true,"with_std":true,"epsilon":1e-12},
476 "counts":[0],
477 "means":[0.0],
478 "m2s":[1.0]
479 }"#;
480 assert!(
481 serde_json::from_str::<StandardScaler>(malformed).is_err(),
482 "count=0 with non-zero m2 must be rejected"
483 );
484 }
485
486 #[cfg(feature = "serde")]
487 #[test]
488 fn scaler_serde_rejects_one_count_nonzero_m2() {
489 let malformed = r#"{
493 "feature_count":1,
494 "config":{"with_mean":true,"with_std":true,"epsilon":1e-12},
495 "counts":[1],
496 "means":[3.0],
497 "m2s":[0.25]
498 }"#;
499 assert!(
500 serde_json::from_str::<StandardScaler>(malformed).is_err(),
501 "count=1 with non-zero m2 must be rejected"
502 );
503 }
504
505 #[cfg(feature = "serde")]
506 #[test]
507 fn scaler_serde_accepts_one_count_finite_mean() {
508 let json = r#"{
512 "feature_count":2,
513 "config":{"with_mean":true,"with_std":true,"epsilon":1e-12},
514 "counts":[1,1],
515 "means":[3.0,-7.5],
516 "m2s":[0.0,0.0]
517 }"#;
518 let scaler: StandardScaler =
519 serde_json::from_str(json).expect("count=1 with finite mean and m2=0 must be accepted");
520 let re = serde_json::to_string(&scaler).unwrap();
522 let _: StandardScaler = serde_json::from_str(&re).unwrap();
523 assert_eq!(scaler.counts, vec![1, 1]);
524 assert_eq!(scaler.means, vec![3.0, -7.5]);
525 assert_eq!(scaler.m2s, vec![0.0, 0.0]);
526 }
527
528 #[cfg(feature = "serde")]
529 #[test]
530 fn scaler_serde_roundtrip_preserves_state() {
531 let mut scaler = StandardScaler::new(2).unwrap();
534 scaler.update(&[1.0, 10.0]).unwrap();
535 scaler.update(&[3.0, 20.0]).unwrap();
536 scaler.update(&[5.0, 30.0]).unwrap();
537 let json = serde_json::to_string(&scaler).unwrap();
538 let restored: StandardScaler = serde_json::from_str(&json).unwrap();
539 assert_eq!(restored.counts, scaler.counts);
540 assert_eq!(restored.means, scaler.means);
541 assert_eq!(restored.m2s, scaler.m2s);
542 let features = [2.0, 15.0];
544 assert_eq!(
545 scaler.transform(&features).unwrap(),
546 restored.transform(&features).unwrap()
547 );
548 }
549
550 #[test]
551 fn transform_hot_path_does_not_scan_invariants() {
552 let mut scaler = StandardScaler::new(3).unwrap();
560 scaler.update(&[1.0, 2.0, 3.0]).unwrap();
561 scaler.update(&[4.0, 5.0, 6.0]).unwrap();
562 let before = scaler.clone();
563 let out = scaler.transform(&[2.5, 3.5, 4.5]).unwrap();
564 assert_eq!(scaler.counts, before.counts);
566 assert_eq!(scaler.means, before.means);
567 assert_eq!(scaler.m2s, before.m2s);
568 assert_eq!(out.len(), 3);
570 }
571
572 #[test]
573 fn constant_feature_uses_scale_one() {
574 let mut s = StandardScaler::new(1).unwrap();
575 for _ in 0..10 {
576 s.update(&[5.0]).unwrap();
577 }
578 let out = s.transform(&[5.0]).unwrap();
580 assert!(out[0].abs() < 1e-12);
581 assert!(!out[0].is_nan());
582 }
583
584 #[test]
585 fn with_mean_false_keeps_offset() {
586 let mut s = StandardScaler::with_config(
587 1,
588 StandardScalerConfig {
589 with_mean: false,
590 with_std: true,
591 epsilon: 1e-12,
592 },
593 )
594 .unwrap();
595 s.update(&[1.0]).unwrap();
596 s.update(&[3.0]).unwrap();
597 let out = s.transform(&[3.0]).unwrap();
599 assert!((out[0] - 3.0).abs() < 1e-9);
600 }
601
602 #[test]
603 fn dimension_mismatch_rejected() {
604 let mut s = StandardScaler::new(3).unwrap();
605 assert!(s.transform(&[1.0, 2.0]).is_err());
606 assert!(s.update(&[1.0, 2.0]).is_err());
607 }
608
609 #[test]
610 fn zero_features_rejected() {
611 assert!(matches!(
612 StandardScaler::new(0),
613 Err(RillError::EmptyFeatures)
614 ));
615 }
616
617 #[test]
618 fn non_finite_rejected() {
619 let mut s = StandardScaler::new(2).unwrap();
620 assert!(s.update(&[1.0, f64::NAN]).is_err());
621 }
622
623 #[test]
624 fn reset_clears_state() {
625 let mut s = StandardScaler::new(1).unwrap();
626 s.update(&[1.0]).unwrap();
627 s.update(&[2.0]).unwrap();
628 s.reset();
629 assert_eq!(s.counts[0], 0);
630 assert_eq!(s.means()[0], 0.0);
631 }
632
633 #[test]
634 fn transform_into_matches_transform_output() {
635 let mut scaler = StandardScaler::new(4).unwrap();
636 scaler.update(&[1.0, 10.0, 100.0, 1000.0]).unwrap();
638 scaler.update(&[3.0, 20.0, 300.0, 3000.0]).unwrap();
639 scaler.update(&[5.0, 30.0, 500.0, 5000.0]).unwrap();
640 let features = [2.0, 15.0, 200.0, 2000.0];
641 let via_transform = scaler.transform(&features).unwrap();
642 let mut via_into = Vec::new();
643 scaler.transform_into(&features, &mut via_into).unwrap();
644 assert_eq!(via_transform, via_into);
645 }
646
647 #[test]
648 fn transform_into_reuses_buffer_capacity() {
649 let mut scaler = StandardScaler::new(3).unwrap();
650 scaler.update(&[1.0, 2.0, 3.0]).unwrap();
651 scaler.update(&[4.0, 5.0, 6.0]).unwrap();
652 let features = [2.5, 3.5, 4.5];
653 let mut buffer = Vec::with_capacity(64);
654 buffer.extend_from_slice(&[-1.0, -2.0, -3.0, -4.0]);
656 scaler.transform_into(&features, &mut buffer).unwrap();
657 assert_eq!(buffer.len(), 3);
658 assert!(buffer.capacity() >= 64);
660 assert_eq!(buffer, scaler.transform(&features).unwrap());
662 }
663
664 #[test]
665 fn transform_into_rejects_dimension_mismatch() {
666 let scaler = StandardScaler::new(3).unwrap();
667 let mut buffer = Vec::new();
668 assert!(scaler.transform_into(&[1.0, 2.0], &mut buffer).is_err());
669 assert!(buffer.is_empty());
671 }
672
673 #[test]
674 fn transform_into_rejects_non_finite_output() {
675 let scaler = StandardScaler::with_config(
678 1,
679 StandardScalerConfig {
680 with_mean: false,
681 with_std: false,
682 epsilon: 1e-12,
683 },
684 )
685 .unwrap();
686 let mut buffer = Vec::new();
687 assert!(scaler.transform_into(&[f64::NAN], &mut buffer).is_err());
688 }
689
690 #[test]
691 fn transform_into_with_zero_state_returns_original() {
692 let scaler = StandardScaler::new(3).unwrap();
693 let features = [1.5, 2.5, 3.5];
694 let mut buffer = Vec::new();
695 scaler.transform_into(&features, &mut buffer).unwrap();
696 assert_eq!(buffer, vec![1.5, 2.5, 3.5]);
698 }
699}