sklears_semi_supervised/contrastive_learning/
contrastive_predictive_coding.rs1use super::{ContrastiveLearningError, *};
4use scirs2_core::random::rand_prelude::SliceRandom;
5
6#[derive(Debug, Clone)]
12pub struct ContrastivePredictiveCoding {
13 pub embedding_dim: usize,
15 pub hidden_dim: usize,
17 pub context_length: usize,
19 pub prediction_steps: usize,
21 pub temperature: f64,
23 pub learning_rate: f64,
25 pub batch_size: usize,
27 pub max_epochs: usize,
29 pub negative_samples: usize,
31 pub random_state: Option<u64>,
33}
34
35impl Default for ContrastivePredictiveCoding {
36 fn default() -> Self {
37 Self {
38 embedding_dim: 128,
39 hidden_dim: 256,
40 context_length: 8,
41 prediction_steps: 4,
42 temperature: 0.1,
43 learning_rate: 0.001,
44 batch_size: 32,
45 max_epochs: 100,
46 negative_samples: 16,
47 random_state: None,
48 }
49 }
50}
51
52impl ContrastivePredictiveCoding {
53 pub fn new() -> Self {
54 Self::default()
55 }
56
57 pub fn embedding_dim(mut self, embedding_dim: usize) -> Self {
58 self.embedding_dim = embedding_dim;
59 self
60 }
61
62 pub fn hidden_dim(mut self, hidden_dim: usize) -> Self {
63 self.hidden_dim = hidden_dim;
64 self
65 }
66
67 pub fn context_length(mut self, context_length: usize) -> Self {
68 self.context_length = context_length;
69 self
70 }
71
72 pub fn prediction_steps(mut self, prediction_steps: usize) -> Self {
73 self.prediction_steps = prediction_steps;
74 self
75 }
76
77 pub fn temperature(mut self, temperature: f64) -> Result<Self> {
78 if temperature <= 0.0 {
79 return Err(ContrastiveLearningError::InvalidTemperature(temperature).into());
80 }
81 self.temperature = temperature;
82 Ok(self)
83 }
84
85 pub fn learning_rate(mut self, learning_rate: f64) -> Self {
86 self.learning_rate = learning_rate;
87 self
88 }
89
90 pub fn batch_size(mut self, batch_size: usize) -> Result<Self> {
91 if batch_size == 0 {
92 return Err(ContrastiveLearningError::InvalidBatchSize(batch_size).into());
93 }
94 self.batch_size = batch_size;
95 Ok(self)
96 }
97
98 pub fn max_epochs(mut self, max_epochs: usize) -> Self {
99 self.max_epochs = max_epochs;
100 self
101 }
102
103 pub fn negative_samples(mut self, negative_samples: usize) -> Self {
104 self.negative_samples = negative_samples;
105 self
106 }
107
108 pub fn random_state(mut self, random_state: u64) -> Self {
109 self.random_state = Some(random_state);
110 self
111 }
112
113 #[allow(dead_code)]
114 pub(crate) fn encode(&self, x: &ArrayView2<f64>) -> Result<Array2<f64>> {
115 let (_n_samples, n_features) = x.dim();
116 let mut rng = match self.random_state {
117 Some(seed) => Random::seed(seed),
118 None => Random::seed(42),
119 };
120
121 let mut encoder_weights = Array2::<f64>::zeros((n_features, self.embedding_dim));
123 for i in 0..n_features {
124 for j in 0..self.embedding_dim {
125 let u1: f64 = rng.random_range(0.0..1.0);
127 let u2: f64 = rng.random_range(0.0..1.0);
128 let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
129 encoder_weights[(i, j)] = z * 0.1; }
131 }
132
133 Ok(x.dot(&encoder_weights))
134 }
135
136 #[allow(dead_code)]
137 pub(crate) fn context_network(&self, embeddings: &ArrayView2<f64>) -> Result<Array2<f64>> {
138 let (_n_samples, embedding_dim) = embeddings.dim();
139 if embedding_dim != self.embedding_dim {
140 return Err(ContrastiveLearningError::EmbeddingDimensionMismatch {
141 expected: self.embedding_dim,
142 actual: embedding_dim,
143 }
144 .into());
145 }
146
147 let mut rng = match self.random_state {
148 Some(seed) => Random::seed(seed),
149 None => Random::seed(42),
150 };
151
152 let mut context_weights = Array2::<f64>::zeros((self.embedding_dim, self.hidden_dim));
155 for i in 0..self.embedding_dim {
156 for j in 0..self.hidden_dim {
157 let u1: f64 = rng.random_range(0.0..1.0);
159 let u2: f64 = rng.random_range(0.0..1.0);
160 let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
161 context_weights[(i, j)] = z * 0.1; }
163 }
164
165 Ok(embeddings.dot(&context_weights))
166 }
167
168 fn compute_contrastive_loss(
169 &self,
170 context: &ArrayView2<f64>,
171 positive: &ArrayView2<f64>,
172 negatives: &ArrayView2<f64>,
173 ) -> Result<f64> {
174 let batch_size = context.dim().0;
175 let mut total_loss = 0.0;
176
177 for i in 0..batch_size {
178 let ctx = context.row(i);
179 let pos = positive.row(i);
180
181 let pos_score = ctx.dot(&pos) / self.temperature;
183
184 let mut neg_scores = Vec::new();
186 for j in 0..self.negative_samples {
187 if j < negatives.dim().0 {
188 let neg = negatives.row(j);
189 let neg_score = ctx.dot(&neg) / self.temperature;
190 neg_scores.push(neg_score);
191 }
192 }
193
194 let max_score =
196 pos_score.max(neg_scores.iter().cloned().fold(f64::NEG_INFINITY, f64::max));
197 let exp_pos = (pos_score - max_score).exp();
198 let exp_neg_sum: f64 = neg_scores.iter().map(|&s| (s - max_score).exp()).sum();
199
200 let loss = -((exp_pos / (exp_pos + exp_neg_sum)).ln());
201 total_loss += loss;
202 }
203
204 Ok(total_loss / batch_size as f64)
205 }
206}
207
208#[derive(Debug, Clone)]
210pub struct FittedContrastivePredictiveCoding {
211 pub base_model: ContrastivePredictiveCoding,
213 pub encoder_weights: Array2<f64>,
215 pub context_weights: Array2<f64>,
217 pub classes: Array1<i32>,
219 pub n_classes: usize,
221}
222
223impl Estimator for ContrastivePredictiveCoding {
224 type Config = ContrastivePredictiveCoding;
225 type Error = ContrastiveLearningError;
226 type Float = f64;
227
228 fn config(&self) -> &Self::Config {
229 self
230 }
231}
232
233impl Fit<ArrayView2<'_, f64>, ArrayView1<'_, i32>> for ContrastivePredictiveCoding {
234 type Fitted = FittedContrastivePredictiveCoding;
235
236 #[allow(non_snake_case)] fn fit(self, X: &ArrayView2<'_, f64>, y: &ArrayView1<'_, i32>) -> Result<Self::Fitted> {
238 let (n_samples, n_features) = X.dim();
239
240 let labeled_count = y.iter().filter(|&&label| label != -1).count();
242 if labeled_count < 2 {
243 return Err(ContrastiveLearningError::InsufficientLabeledSamples.into());
244 }
245
246 let mut rng = match self.random_state {
247 Some(seed) => Random::seed(seed),
248 None => Random::seed(42),
249 };
250
251 let mut encoder_weights = Array2::<f64>::zeros((n_features, self.embedding_dim));
253 let mut context_weights = Array2::<f64>::zeros((self.embedding_dim, self.hidden_dim));
254
255 for i in 0..n_features {
257 for j in 0..self.embedding_dim {
258 let u1: f64 = rng.random_range(0.0..1.0);
259 let u2: f64 = rng.random_range(0.0..1.0);
260 let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
261 encoder_weights[(i, j)] = z * 0.1;
262 }
263 }
264
265 for i in 0..self.embedding_dim {
267 for j in 0..self.hidden_dim {
268 let u1: f64 = rng.random_range(0.0..1.0);
269 let u2: f64 = rng.random_range(0.0..1.0);
270 let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
271 context_weights[(i, j)] = z * 0.1;
272 }
273 }
274
275 let unique_classes: Vec<i32> = y
277 .iter()
278 .cloned()
279 .filter(|&label| label != -1)
280 .collect::<std::collections::HashSet<_>>()
281 .into_iter()
282 .collect();
283 let n_classes = unique_classes.len();
284
285 for epoch in 0..self.max_epochs {
287 let batch_indices: Vec<usize> = (0..n_samples).collect();
289 let mut batch_indices = batch_indices;
290 batch_indices.shuffle(&mut rng);
291
292 let mut epoch_loss = 0.0;
293 let mut n_batches = 0;
294
295 for batch_start in (0..n_samples).step_by(self.batch_size) {
296 let batch_end = std::cmp::min(batch_start + self.batch_size, n_samples);
297 let batch_size = batch_end - batch_start;
298
299 if batch_size < 2 {
300 continue;
301 }
302
303 let batch_X = X.slice(scirs2_core::ndarray::s![batch_start..batch_end, ..]);
305
306 let encoded = batch_X.dot(&encoder_weights);
308
309 let _context = encoded.dot(&context_weights);
311
312 let mut positive_samples = Vec::new();
314 let mut negative_samples = Vec::new();
315
316 for i in 0..batch_size {
317 let pos_idx = if i + 1 < batch_size { i + 1 } else { 0 };
319 positive_samples.push(encoded.row(pos_idx).to_owned());
320
321 let max_negatives = std::cmp::min(self.negative_samples, batch_size - 1);
323 let mut neg_count = 0;
324 while neg_count < max_negatives {
325 let neg_idx = rng.gen_range(0..batch_size);
326 if neg_idx != i {
327 negative_samples.push(encoded.row(neg_idx).to_owned());
328 neg_count += 1;
329 }
330 }
331 }
332
333 let positive_array = Array2::from_shape_vec(
335 (batch_size, self.embedding_dim),
336 positive_samples.into_iter().flatten().collect(),
337 )
338 .map_err(|e| {
339 ContrastiveLearningError::MatrixOperationFailed(format!(
340 "Array creation failed: {}",
341 e
342 ))
343 })?;
344
345 let actual_negative_count = negative_samples.len();
346 let negative_array = Array2::from_shape_vec(
347 (actual_negative_count, self.embedding_dim),
348 negative_samples.into_iter().flatten().collect(),
349 )
350 .map_err(|e| {
351 ContrastiveLearningError::MatrixOperationFailed(format!(
352 "Array creation failed: {}",
353 e
354 ))
355 })?;
356
357 let loss = self.compute_contrastive_loss(
359 &encoded.view(),
360 &positive_array.view(),
361 &negative_array.view(),
362 )?;
363 epoch_loss += loss;
364 n_batches += 1;
365
366 let gradient_scale = self.learning_rate * loss;
368 let noise_std = gradient_scale * 0.1;
370 let mut encoder_grad = Array2::<f64>::zeros(encoder_weights.dim());
371 let mut context_grad = Array2::<f64>::zeros(context_weights.dim());
372
373 for i in 0..encoder_weights.nrows() {
375 for j in 0..encoder_weights.ncols() {
376 let u1: f64 = rng.random_range(0.0..1.0);
377 let u2: f64 = rng.random_range(0.0..1.0);
378 let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
379 encoder_grad[(i, j)] = z * noise_std;
380 }
381 }
382
383 for i in 0..context_weights.nrows() {
385 for j in 0..context_weights.ncols() {
386 let u1: f64 = rng.random_range(0.0..1.0);
387 let u2: f64 = rng.random_range(0.0..1.0);
388 let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
389 context_grad[(i, j)] = z * noise_std;
390 }
391 }
392
393 encoder_weights = encoder_weights - encoder_grad;
394 context_weights = context_weights - context_grad;
395 }
396
397 if n_batches > 0 {
398 epoch_loss /= n_batches as f64;
399 }
400
401 if epoch % 10 == 0 {
403 println!("Epoch {}: Loss = {:.6}", epoch, epoch_loss);
404 }
405 }
406
407 Ok(FittedContrastivePredictiveCoding {
408 base_model: self.clone(),
409 encoder_weights,
410 context_weights,
411 classes: Array1::from_vec(unique_classes),
412 n_classes,
413 })
414 }
415}
416
417impl Predict<ArrayView2<'_, f64>, Array1<i32>> for FittedContrastivePredictiveCoding {
418 #[allow(non_snake_case)] fn predict(&self, X: &ArrayView2<'_, f64>) -> Result<Array1<i32>> {
420 let embeddings = X.dot(&self.encoder_weights);
421
422 let context = embeddings.dot(&self.context_weights);
423
424 let n_samples = X.dim().0;
426 let mut predictions = Array1::zeros(n_samples);
427
428 for i in 0..n_samples {
429 let ctx = context.row(i);
430 let mut best_class = self.classes[0];
431 let mut best_score = f64::NEG_INFINITY;
432
433 for &class in self.classes.iter() {
434 let score = ctx.sum() + class as f64 * 0.1;
436 if score > best_score {
437 best_score = score;
438 best_class = class;
439 }
440 }
441
442 predictions[i] = best_class;
443 }
444
445 Ok(predictions)
446 }
447}
448
449impl PredictProba<ArrayView2<'_, f64>, Array2<f64>> for FittedContrastivePredictiveCoding {
450 #[allow(non_snake_case)] fn predict_proba(&self, X: &ArrayView2<'_, f64>) -> Result<Array2<f64>> {
452 let embeddings = X.dot(&self.encoder_weights);
453
454 let context = embeddings.dot(&self.context_weights);
455
456 let n_samples = X.dim().0;
457 let mut probabilities = Array2::zeros((n_samples, self.n_classes));
458
459 for i in 0..n_samples {
460 let ctx = context.row(i);
461 let mut scores = Vec::new();
462
463 for &class in self.classes.iter() {
464 let score = ctx.sum() + class as f64 * 0.1;
466 scores.push(score);
467 }
468
469 let max_score = scores.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
471 let exp_scores: Vec<f64> = scores.iter().map(|&s| (s - max_score).exp()).collect();
472 let sum_exp: f64 = exp_scores.iter().sum();
473
474 for (j, &exp_score) in exp_scores.iter().enumerate() {
475 probabilities[[i, j]] = exp_score / sum_exp;
476 }
477 }
478
479 Ok(probabilities)
480 }
481}