1use imbalance_rs::{RandomOverSampler as ImbRos, Sampler, Smote as ImbSmote};
12use ndarray::{Array1, Array2};
13
14use crate::error::{Error, Result};
15use crate::frame::Frame;
16use crate::traits::Balancer;
17
18fn frame_to_array2(frame: &Frame) -> Result<Array2<f64>> {
19 let (n, p) = frame.shape();
20 Array2::from_shape_vec((n, p), frame.buf().to_vec())
21 .map_err(|e| Error::Backend(format!("ndarray conversion failed: {e}")))
22}
23
24fn array2_to_frame(arr: &Array2<f64>, columns: &[String]) -> Result<Frame> {
25 let (n, p) = arr.dim();
26 let buf: Vec<f64> = arr.iter().copied().collect();
29 Frame::new(buf, n, p, columns.to_vec())
30}
31
32#[derive(Clone, Debug)]
37pub struct Smote {
38 k_neighbors: usize,
39 seed: u64,
40}
41
42impl Smote {
43 pub fn new() -> Self {
45 Smote {
46 k_neighbors: 5,
47 seed: 0,
48 }
49 }
50
51 pub fn k_neighbors(mut self, k: usize) -> Self {
53 self.k_neighbors = k;
54 self
55 }
56
57 pub fn random_state(mut self, seed: u64) -> Self {
59 self.seed = seed;
60 self
61 }
62}
63
64impl Default for Smote {
65 fn default() -> Self {
66 Smote::new()
67 }
68}
69
70impl Balancer for Smote {
71 fn name(&self) -> &'static str {
72 "Smote"
73 }
74
75 fn fit_resample(&self, features: &Frame, target: &[f64]) -> Result<(Frame, Vec<f64>)> {
76 let x = frame_to_array2(features)?;
77 let y: Array1<i64> =
78 Array1::from(target.iter().map(|v| v.round() as i64).collect::<Vec<_>>());
79 let sampler = ImbSmote::new()
80 .k_neighbors(self.k_neighbors)
81 .random_state(self.seed);
82 let (xr, yr) = sampler
83 .fit_resample(&x, &y)
84 .map_err(|e| Error::Backend(format!("SMOTE failed: {e}")))?;
85 let frame = array2_to_frame(&xr, features.columns())?;
86 let target = yr.iter().map(|&l| l as f64).collect();
87 Ok((frame, target))
88 }
89}
90
91#[derive(Clone, Debug, Default)]
95pub struct RandomOverSampler;
96
97impl RandomOverSampler {
98 pub fn new() -> Self {
99 RandomOverSampler
100 }
101}
102
103impl Balancer for RandomOverSampler {
104 fn name(&self) -> &'static str {
105 "RandomOverSampler"
106 }
107
108 fn fit_resample(&self, features: &Frame, target: &[f64]) -> Result<(Frame, Vec<f64>)> {
109 let x = frame_to_array2(features)?;
110 let y: Array1<i64> =
111 Array1::from(target.iter().map(|v| v.round() as i64).collect::<Vec<_>>());
112 let (xr, yr) = ImbRos::new()
113 .fit_resample(&x, &y)
114 .map_err(|e| Error::Backend(format!("RandomOverSampler failed: {e}")))?;
115 let frame = array2_to_frame(&xr, features.columns())?;
116 let target = yr.iter().map(|&l| l as f64).collect();
117 Ok((frame, target))
118 }
119}
120
121#[cfg(test)]
122mod tests {
123 use super::*;
124
125 #[test]
126 fn smote_balances_the_minority_class() {
127 let x = Frame::from_rows(
129 vec![
130 vec![0.0, 0.0],
131 vec![0.1, 0.2],
132 vec![0.2, 0.1],
133 vec![0.3, 0.0],
134 vec![0.0, 0.3],
135 vec![0.2, 0.2],
136 vec![9.0, 9.0],
137 vec![9.1, 9.2],
138 ],
139 vec!["a".into(), "b".into()],
140 )
141 .unwrap();
142 let y = vec![0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0, 1.0];
143
144 let (xr, yr) = Smote::new()
145 .k_neighbors(1)
146 .random_state(7)
147 .fit_resample(&x, &y)
148 .unwrap();
149 let zeros = yr.iter().filter(|&&v| v == 0.0).count();
150 let ones = yr.iter().filter(|&&v| v == 1.0).count();
151 assert_eq!(zeros, 6);
152 assert_eq!(ones, 6, "minority class should be oversampled to parity");
153 assert_eq!(xr.nrows(), yr.len());
154 assert_eq!(xr.ncols(), 2);
155 }
156}