Skip to main content

sklears_semi_supervised/contrastive_learning/
contrastive_predictive_coding.rs

1//! Contrastive Predictive Coding (CPC) implementation for semi-supervised learning
2
3use super::{ContrastiveLearningError, *};
4use scirs2_core::random::rand_prelude::SliceRandom;
5
6/// Contrastive Predictive Coding (CPC) for semi-supervised learning
7///
8/// CPC learns representations by predicting future observations from past contexts
9/// in a contrastive manner. It maximizes mutual information between contexts and
10/// positive samples while minimizing it for negative samples.
11#[derive(Debug, Clone)]
12pub struct ContrastivePredictiveCoding {
13    /// embedding_dim
14    pub embedding_dim: usize,
15    /// hidden_dim
16    pub hidden_dim: usize,
17    /// context_length
18    pub context_length: usize,
19    /// prediction_steps
20    pub prediction_steps: usize,
21    /// temperature
22    pub temperature: f64,
23    /// learning_rate
24    pub learning_rate: f64,
25    /// batch_size
26    pub batch_size: usize,
27    /// max_epochs
28    pub max_epochs: usize,
29    /// negative_samples
30    pub negative_samples: usize,
31    /// random_state
32    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        // Simple linear encoder for demonstration - create weights manually
122        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                // Generate normal distributed random number using Box-Muller transform
126                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; // mean=0.0, std=0.1
130            }
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        // Simple context network (could be LSTM/GRU in practice)
153        // Create context weights manually
154        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                // Generate normal distributed random number using Box-Muller transform
158                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; // mean=0.0, std=0.1
162            }
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            // Compute positive score
182            let pos_score = ctx.dot(&pos) / self.temperature;
183
184            // Compute negative scores
185            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            // Compute softmax loss
195            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/// Fitted Contrastive Predictive Coding model
209#[derive(Debug, Clone)]
210pub struct FittedContrastivePredictiveCoding {
211    /// base_model
212    pub base_model: ContrastivePredictiveCoding,
213    /// encoder_weights
214    pub encoder_weights: Array2<f64>,
215    /// context_weights
216    pub context_weights: Array2<f64>,
217    /// classes
218    pub classes: Array1<i32>,
219    /// n_classes
220    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)] // standard ML notation
237    fn fit(self, X: &ArrayView2<'_, f64>, y: &ArrayView1<'_, i32>) -> Result<Self::Fitted> {
238        let (n_samples, n_features) = X.dim();
239
240        // Check for sufficient labeled samples
241        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        // Initialize encoder and context networks
252        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        // Fill encoder weights with normal distribution (mean=0.0, std=0.1)
256        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        // Fill context weights with normal distribution (mean=0.0, std=0.1)
266        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        // Get unique classes
276        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        // Training loop
286        for epoch in 0..self.max_epochs {
287            // Generate batches
288            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                // Get batch data
304                let batch_X = X.slice(scirs2_core::ndarray::s![batch_start..batch_end, ..]);
305
306                // Encode batch
307                let encoded = batch_X.dot(&encoder_weights);
308
309                // Context network
310                let _context = encoded.dot(&context_weights);
311
312                // Generate positive and negative samples
313                let mut positive_samples = Vec::new();
314                let mut negative_samples = Vec::new();
315
316                for i in 0..batch_size {
317                    // Use next sample as positive (temporal structure)
318                    let pos_idx = if i + 1 < batch_size { i + 1 } else { 0 };
319                    positive_samples.push(encoded.row(pos_idx).to_owned());
320
321                    // Random negative samples
322                    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                // Convert to arrays
334                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                // Compute loss using encoded representations
358                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                // Simple gradient update (in practice, would use proper backpropagation)
367                let gradient_scale = self.learning_rate * loss;
368                // Create gradient noise manually
369                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                // Fill encoder grad with normal noise
374                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                // Fill context grad with normal noise
384                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            // Early stopping or convergence check could be added here
402            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)] // standard ML notation
419    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        // Simple nearest class prediction based on context representations
425        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                // Simple scoring based on context magnitude (placeholder)
435                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)] // standard ML notation
451    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                // Simple scoring based on context (placeholder)
465                let score = ctx.sum() + class as f64 * 0.1;
466                scores.push(score);
467            }
468
469            // Softmax normalization
470            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}