sklears_semi_supervised/batch_active_learning/
core_set.rs1use super::{BatchActiveLearningError, *};
4
5#[derive(Debug, Clone)]
10pub struct CoreSetApproach {
11 pub batch_size: usize,
13 pub distance_metric: String,
15 pub initialization: String,
17 pub max_iter: usize,
19 pub random_state: Option<u64>,
21}
22
23impl Default for CoreSetApproach {
24 fn default() -> Self {
25 Self {
26 batch_size: 10,
27 distance_metric: "euclidean".to_string(),
28 initialization: "farthest_first".to_string(),
29 max_iter: 100,
30 random_state: None,
31 }
32 }
33}
34
35impl CoreSetApproach {
36 pub fn new() -> Self {
37 Self::default()
38 }
39
40 pub fn batch_size(mut self, batch_size: usize) -> Result<Self> {
41 if batch_size == 0 {
42 return Err(BatchActiveLearningError::InvalidBatchSize(batch_size).into());
43 }
44 self.batch_size = batch_size;
45 Ok(self)
46 }
47
48 pub fn distance_metric(mut self, distance_metric: String) -> Self {
49 self.distance_metric = distance_metric;
50 self
51 }
52
53 pub fn initialization(mut self, initialization: String) -> Self {
54 self.initialization = initialization;
55 self
56 }
57
58 pub fn max_iter(mut self, max_iter: usize) -> Self {
59 self.max_iter = max_iter;
60 self
61 }
62
63 pub fn random_state(mut self, random_state: u64) -> Self {
64 self.random_state = Some(random_state);
65 self
66 }
67
68 fn compute_distance(&self, x1: &ArrayView1<f64>, x2: &ArrayView1<f64>) -> Result<f64> {
69 match self.distance_metric.as_str() {
70 "euclidean" => {
71 let dist = x1
72 .iter()
73 .zip(x2.iter())
74 .map(|(a, b)| (a - b).powi(2))
75 .sum::<f64>()
76 .sqrt();
77 Ok(dist)
78 }
79 "manhattan" => {
80 let dist = x1
81 .iter()
82 .zip(x2.iter())
83 .map(|(a, b)| (a - b).abs())
84 .sum::<f64>();
85 Ok(dist)
86 }
87 _ => Err(
88 BatchActiveLearningError::InvalidDistanceMetric(self.distance_metric.clone())
89 .into(),
90 ),
91 }
92 }
93
94 #[allow(non_snake_case)] fn farthest_first_initialization(&self, X: &ArrayView2<f64>) -> Result<Vec<usize>> {
96 let n_samples = X.dim().0;
97 let mut rng = match self.random_state {
98 Some(seed) => Random::seed(seed),
99 None => Random::seed(42),
100 };
101
102 if n_samples < self.batch_size {
103 return Err(BatchActiveLearningError::InsufficientUnlabeledSamples.into());
104 }
105
106 let mut selected_indices = Vec::new();
107 let mut distances = vec![f64::INFINITY; n_samples];
108
109 let first_idx = rng.gen_range(0..n_samples);
111 selected_indices.push(first_idx);
112
113 for (i, dist) in distances.iter_mut().enumerate() {
115 if i != first_idx {
116 *dist = self.compute_distance(&X.row(i), &X.row(first_idx))?;
117 }
118 }
119
120 for _ in 1..self.batch_size {
122 let mut max_distance = 0.0;
124 let mut best_idx = 0;
125
126 for (i, &dist) in distances.iter().enumerate() {
127 if !selected_indices.contains(&i) && dist > max_distance {
128 max_distance = dist;
129 best_idx = i;
130 }
131 }
132
133 selected_indices.push(best_idx);
134
135 for (i, dist) in distances.iter_mut().enumerate() {
137 if !selected_indices.contains(&i) {
138 let new_distance = self.compute_distance(&X.row(i), &X.row(best_idx))?;
139 *dist = (*dist).min(new_distance);
140 }
141 }
142 }
143
144 Ok(selected_indices)
145 }
146
147 #[allow(non_snake_case)] fn k_center_greedy(&self, X: &ArrayView2<f64>) -> Result<Vec<usize>> {
149 let n_samples = X.dim().0;
150 let mut selected_indices = Vec::new();
151 let mut distances = vec![f64::INFINITY; n_samples];
152
153 if n_samples < self.batch_size {
154 return Err(BatchActiveLearningError::InsufficientUnlabeledSamples.into());
155 }
156
157 let mut centroid = Array1::zeros(X.dim().1);
159 for i in 0..n_samples {
160 centroid = centroid + X.row(i);
161 }
162 centroid /= n_samples as f64;
163
164 let mut min_distance = f64::INFINITY;
166 let mut first_idx = 0;
167 for i in 0..n_samples {
168 let distance = self.compute_distance(&X.row(i), ¢roid.view())?;
169 if distance < min_distance {
170 min_distance = distance;
171 first_idx = i;
172 }
173 }
174
175 selected_indices.push(first_idx);
176
177 for (i, dist) in distances.iter_mut().enumerate() {
179 if i != first_idx {
180 *dist = self.compute_distance(&X.row(i), &X.row(first_idx))?;
181 }
182 }
183
184 for _ in 1..self.batch_size {
186 let mut max_distance = 0.0;
187 let mut best_idx = 0;
188
189 for (i, &dist) in distances.iter().enumerate() {
190 if !selected_indices.contains(&i) && dist > max_distance {
191 max_distance = dist;
192 best_idx = i;
193 }
194 }
195
196 selected_indices.push(best_idx);
197
198 for (i, dist) in distances.iter_mut().enumerate() {
200 if !selected_indices.contains(&i) {
201 let new_distance = self.compute_distance(&X.row(i), &X.row(best_idx))?;
202 *dist = (*dist).min(new_distance);
203 }
204 }
205 }
206
207 Ok(selected_indices)
208 }
209
210 #[allow(non_snake_case)] pub fn query(
212 &self,
213 X: &ArrayView2<f64>,
214 _probabilities: &ArrayView2<f64>,
215 ) -> Result<Vec<usize>> {
216 match self.initialization.as_str() {
217 "farthest_first" => self.farthest_first_initialization(X),
218 "k_center_greedy" => self.k_center_greedy(X),
219 _ => self.farthest_first_initialization(X), }
221 }
222}