Whole-workspace rustfmt pass picked up while iterating on Mamba GPU backward work. Verified formatting-only via diff sampling; no logic changed. Co-Authored-By: Claude Sonnet 5 <[email protected]>
788 lines
24 KiB
Rust
788 lines
24 KiB
Rust
//! CPC (Contrastive Predictive Coding) Implementation
|
|
//!
|
|
//! Self-supervised learning method that learns representations by predicting future
|
|
//! representations in latent space using autoregressive models.
|
|
//!
|
|
//! Based on "Representation Learning with Contrastive Predictive Coding" (van den Oord et al., 2018)
|
|
//!
|
|
//! ## Key Algorithm Components
|
|
//!
|
|
//! 1. **Encoder network**: g_enc(x_t) → z_t for extracting representations
|
|
//! 2. **Autoregressive context**: g_ar(z_≤t) → c_t using GRU/Transformer
|
|
//! 3. **Future prediction**: Predict z_{t+k} from c_t for k steps ahead
|
|
//! 4. **Contrastive loss**: InfoNCE to distinguish positive from negative samples
|
|
//! 5. **Score function**: f_k(c_t, z_{t+k}) = exp(z_{t+k}^T W_k c_t)
|
|
//!
|
|
//! ## Modality Support
|
|
//!
|
|
//! - **Images**: Spatial CPC with patches in raster order, masked convolutions
|
|
//! - **Audio**: Temporal CPC with sequence modeling, Wav2Vec-style encoding
|
|
//!
|
|
//! ## Features
|
|
//!
|
|
//! - CNN and Wav2Vec encoders for different modalities
|
|
//! - GRU and Transformer context networks
|
|
//! - Multiple prediction steps with separate heads
|
|
//! - InfoNCE loss with flexible negative sampling
|
|
//! - Batch and memory bank negative sampling strategies
|
|
|
|
use crate::prelude::*;
|
|
use parking_lot::RwLock;
|
|
use std::collections::VecDeque;
|
|
use std::sync::Arc;
|
|
|
|
/// CPC configuration parameters
|
|
#[derive(Debug, Clone)]
|
|
pub struct CPCConfig {
|
|
pub encoder_dim: usize,
|
|
pub context_dim: usize,
|
|
pub num_pred_steps: usize,
|
|
pub context_network: ContextNetworkType,
|
|
pub hidden_dim: usize,
|
|
pub temperature: f32,
|
|
pub negative_samples: usize,
|
|
pub encoder_type: EncoderType,
|
|
}
|
|
|
|
impl CPCConfig {
|
|
pub fn new(encoder_dim: usize, context_dim: usize) -> Self {
|
|
Self {
|
|
encoder_dim,
|
|
context_dim,
|
|
num_pred_steps: 4,
|
|
context_network: ContextNetworkType::GRU,
|
|
hidden_dim: 256,
|
|
temperature: 0.07,
|
|
negative_samples: 16,
|
|
encoder_type: EncoderType::CNN,
|
|
}
|
|
}
|
|
|
|
pub fn with_num_pred_steps(mut self, num_pred_steps: usize) -> Self {
|
|
self.num_pred_steps = num_pred_steps;
|
|
self
|
|
}
|
|
|
|
pub fn with_context_network(mut self, context_network: ContextNetworkType) -> Self {
|
|
self.context_network = context_network;
|
|
self
|
|
}
|
|
|
|
pub fn with_hidden_dim(mut self, hidden_dim: usize) -> Self {
|
|
self.hidden_dim = hidden_dim;
|
|
self
|
|
}
|
|
|
|
pub fn with_temperature(mut self, temperature: f32) -> Self {
|
|
self.temperature = temperature;
|
|
self
|
|
}
|
|
|
|
pub fn with_negative_samples(mut self, negative_samples: usize) -> Self {
|
|
self.negative_samples = negative_samples;
|
|
self
|
|
}
|
|
|
|
pub fn with_encoder_type(mut self, encoder_type: EncoderType) -> Self {
|
|
self.encoder_type = encoder_type;
|
|
self
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
pub enum ContextNetworkType {
|
|
GRU,
|
|
Transformer,
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
pub enum EncoderType {
|
|
CNN,
|
|
Wav2Vec,
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct EncoderConfig {
|
|
pub input_channels: usize,
|
|
pub output_dim: usize,
|
|
pub encoder_type: EncoderType,
|
|
}
|
|
|
|
pub struct CNNEncoder {
|
|
conv1: Tensor,
|
|
conv1_bias: Tensor,
|
|
conv2: Tensor,
|
|
conv2_bias: Tensor,
|
|
conv3: Tensor,
|
|
conv3_bias: Tensor,
|
|
proj: Tensor,
|
|
proj_bias: Tensor,
|
|
pub config: EncoderConfig,
|
|
device: Device,
|
|
}
|
|
|
|
impl CNNEncoder {
|
|
pub fn new(config: EncoderConfig, device: &Device) -> Result<Self> {
|
|
let input_channels = config.input_channels;
|
|
let hidden_channels = 64;
|
|
let output_dim = config.output_dim;
|
|
|
|
let conv1 = Tensor::randn(&[hidden_channels, input_channels, 3, 3], device)?.mul(
|
|
&Tensor::full(&[], (2.0 / (input_channels * 9) as f32).sqrt(), device)?,
|
|
)?;
|
|
let conv1_bias = Tensor::zeros(&[hidden_channels], device)?;
|
|
|
|
let conv2 = Tensor::randn(&[hidden_channels * 2, hidden_channels, 3, 3], device)?.mul(
|
|
&Tensor::full(&[], (2.0 / (hidden_channels * 9) as f32).sqrt(), device)?,
|
|
)?;
|
|
let conv2_bias = Tensor::zeros(&[hidden_channels * 2], device)?;
|
|
|
|
let conv3 = Tensor::randn(&[output_dim, hidden_channels * 2, 3, 3], device)?.mul(
|
|
&Tensor::full(&[], (2.0 / (hidden_channels * 2 * 9) as f32).sqrt(), device)?,
|
|
)?;
|
|
let conv3_bias = Tensor::zeros(&[output_dim], device)?;
|
|
|
|
let proj = Tensor::randn(&[output_dim, output_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
(1.0 / output_dim as f32).sqrt(),
|
|
device,
|
|
)?)?;
|
|
let proj_bias = Tensor::zeros(&[output_dim], device)?;
|
|
|
|
Ok(Self {
|
|
conv1,
|
|
conv1_bias,
|
|
conv2,
|
|
conv2_bias,
|
|
conv3,
|
|
conv3_bias,
|
|
proj,
|
|
proj_bias,
|
|
config,
|
|
device: device.clone(),
|
|
})
|
|
}
|
|
|
|
pub fn forward(&self, input: &Tensor) -> Result<Tensor> {
|
|
let batch_size = input.shape()[0];
|
|
|
|
let h1 = self
|
|
.apply_conv2d(input, &self.conv1, &self.conv1_bias, 1)?
|
|
.relu()?;
|
|
let h2 = self
|
|
.apply_conv2d(&h1, &self.conv2, &self.conv2_bias, 2)?
|
|
.relu()?;
|
|
let h3 = self
|
|
.apply_conv2d(&h2, &self.conv3, &self.conv3_bias, 1)?
|
|
.relu()?;
|
|
|
|
let out_shape = h3.shape();
|
|
let (height, width) = (out_shape[2], out_shape[3]);
|
|
|
|
let features = h3.permute(&[0, 2, 3, 1])?.reshape(&[
|
|
batch_size,
|
|
height * width,
|
|
self.config.output_dim,
|
|
])?;
|
|
|
|
Ok(features.matmul(&self.proj)?.add(&self.proj_bias)?)
|
|
}
|
|
|
|
fn apply_conv2d(
|
|
&self,
|
|
input: &Tensor,
|
|
weight: &Tensor,
|
|
bias: &Tensor,
|
|
stride: usize,
|
|
) -> Result<Tensor> {
|
|
let input_shape = input.shape();
|
|
let weight_shape = weight.shape();
|
|
let (batch_size, in_channels) = (input_shape[0], input_shape[1]);
|
|
let (in_height, in_width) = (input_shape[2], input_shape[3]);
|
|
let (out_channels, kernel_size) = (weight_shape[0], weight_shape[2]);
|
|
|
|
let (out_height, out_width) = (
|
|
(in_height - kernel_size) / stride + 1,
|
|
(in_width - kernel_size) / stride + 1,
|
|
);
|
|
|
|
let resized = if stride == 2 {
|
|
input.narrow(2, 0, out_height)?.narrow(3, 0, out_width)?
|
|
} else {
|
|
input.narrow(2, 0, out_height)?.narrow(3, 0, out_width)?
|
|
};
|
|
|
|
let resized_flat = resized.reshape(&[batch_size, in_channels, out_height * out_width])?;
|
|
let weight_simplified = weight
|
|
.reshape(&[out_channels, in_channels * kernel_size * kernel_size])?
|
|
.narrow(1, 0, in_channels)?;
|
|
|
|
let conv_out = resized_flat
|
|
.transpose(1, 2)?
|
|
.matmul(&weight_simplified.transpose(0, 1)?)?
|
|
.transpose(1, 2)?
|
|
.reshape(&[batch_size, out_channels, out_height, out_width])?;
|
|
|
|
Ok(conv_out.add(&bias.unsqueeze(0)?.unsqueeze(3)?.unsqueeze(4)?)?)
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct ContextNetworkConfig {
|
|
pub input_dim: usize,
|
|
pub hidden_dim: usize,
|
|
pub network_type: ContextNetworkType,
|
|
}
|
|
|
|
/// GRU-based context network for autoregressive modeling
|
|
pub struct GRUContextNetwork {
|
|
reset_ih: Tensor,
|
|
reset_hh: Tensor,
|
|
reset_bias: Tensor,
|
|
update_ih: Tensor,
|
|
update_hh: Tensor,
|
|
update_bias: Tensor,
|
|
new_ih: Tensor,
|
|
new_hh: Tensor,
|
|
new_bias: Tensor,
|
|
pub config: ContextNetworkConfig,
|
|
device: Device,
|
|
}
|
|
|
|
impl GRUContextNetwork {
|
|
pub fn new(config: ContextNetworkConfig, device: &Device) -> Result<Self> {
|
|
let (input_dim, hidden_dim) = (config.input_dim, config.hidden_dim);
|
|
let scale = (1.0 / hidden_dim as f32).sqrt();
|
|
|
|
let reset_ih = Tensor::randn(&[input_dim, hidden_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
scale,
|
|
device,
|
|
)?)?;
|
|
let reset_hh = Tensor::randn(&[hidden_dim, hidden_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
scale,
|
|
device,
|
|
)?)?;
|
|
let reset_bias = Tensor::zeros(&[hidden_dim], device)?;
|
|
|
|
let update_ih = Tensor::randn(&[input_dim, hidden_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
scale,
|
|
device,
|
|
)?)?;
|
|
let update_hh = Tensor::randn(&[hidden_dim, hidden_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
scale,
|
|
device,
|
|
)?)?;
|
|
let update_bias = Tensor::zeros(&[hidden_dim], device)?;
|
|
|
|
let new_ih = Tensor::randn(&[input_dim, hidden_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
scale,
|
|
device,
|
|
)?)?;
|
|
let new_hh = Tensor::randn(&[hidden_dim, hidden_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
scale,
|
|
device,
|
|
)?)?;
|
|
let new_bias = Tensor::zeros(&[hidden_dim], device)?;
|
|
|
|
Ok(Self {
|
|
reset_ih,
|
|
reset_hh,
|
|
reset_bias,
|
|
update_ih,
|
|
update_hh,
|
|
update_bias,
|
|
new_ih,
|
|
new_hh,
|
|
new_bias,
|
|
config,
|
|
device: device.clone(),
|
|
})
|
|
}
|
|
|
|
pub fn forward(&self, input: &Tensor) -> Result<Tensor> {
|
|
let (batch_size, seq_len, hidden_dim) =
|
|
(input.shape()[0], input.shape()[1], self.config.hidden_dim);
|
|
let mut hidden = Tensor::zeros(&[batch_size, hidden_dim], &self.device)?;
|
|
let mut outputs = Vec::new();
|
|
|
|
for t in 0..seq_len {
|
|
let input_t = input.narrow(1, t, 1)?.squeeze(Some(1))?;
|
|
hidden = self.gru_cell(&input_t, &hidden)?;
|
|
outputs.push(hidden.unsqueeze(1)?);
|
|
}
|
|
|
|
Ok(Tensor::cat(&outputs, 1)?)
|
|
}
|
|
|
|
/// Single GRU cell computation
|
|
fn gru_cell(&self, input: &Tensor, hidden: &Tensor) -> Result<Tensor> {
|
|
// Reset gate
|
|
let reset_gate = input
|
|
.matmul(&self.reset_ih)?
|
|
.add(&hidden.matmul(&self.reset_hh)?)?
|
|
.add(&self.reset_bias)?
|
|
.sigmoid()?;
|
|
|
|
// Update gate
|
|
let update_gate = input
|
|
.matmul(&self.update_ih)?
|
|
.add(&hidden.matmul(&self.update_hh)?)?
|
|
.add(&self.update_bias)?
|
|
.sigmoid()?;
|
|
|
|
let reset_hidden = reset_gate.mul(hidden)?;
|
|
let new_gate = input
|
|
.matmul(&self.new_ih)?
|
|
.add(&reset_hidden.matmul(&self.new_hh)?)?
|
|
.add(&self.new_bias)?
|
|
.tanh()?;
|
|
|
|
let one = Tensor::ones_like(&update_gate)?;
|
|
Ok(one
|
|
.sub(&update_gate)?
|
|
.mul(&new_gate)?
|
|
.add(&update_gate.mul(hidden)?)?)
|
|
}
|
|
}
|
|
|
|
/// Configuration for prediction heads
|
|
#[derive(Debug, Clone)]
|
|
pub struct PredictionHeadsConfig {
|
|
/// Context network output dimension
|
|
pub context_dim: usize,
|
|
/// Encoder output dimension to predict
|
|
pub encoder_dim: usize,
|
|
/// Number of prediction steps
|
|
pub num_pred_steps: usize,
|
|
}
|
|
|
|
/// Multiple prediction heads for different future steps
|
|
pub struct PredictionHeads {
|
|
heads: Vec<PredictionHead>,
|
|
config: PredictionHeadsConfig,
|
|
}
|
|
|
|
impl PredictionHeads {
|
|
/// Create new prediction heads
|
|
pub fn new(config: PredictionHeadsConfig, device: &Device) -> Result<Self> {
|
|
let mut heads = Vec::new();
|
|
|
|
for k in 1..=config.num_pred_steps {
|
|
let head = PredictionHead::new(config.context_dim, config.encoder_dim, k, device)?;
|
|
heads.push(head);
|
|
}
|
|
|
|
Ok(Self { heads, config })
|
|
}
|
|
|
|
pub fn forward(&self, context: &Tensor) -> Result<Vec<Tensor>> {
|
|
self.heads
|
|
.iter()
|
|
.map(|head| head.forward(context))
|
|
.collect()
|
|
}
|
|
}
|
|
|
|
pub struct PredictionHead {
|
|
weight: Tensor,
|
|
bias: Tensor,
|
|
step: usize,
|
|
}
|
|
|
|
impl PredictionHead {
|
|
pub fn new(
|
|
context_dim: usize,
|
|
encoder_dim: usize,
|
|
step: usize,
|
|
device: &Device,
|
|
) -> Result<Self> {
|
|
let weight = Tensor::randn(&[context_dim, encoder_dim], device)?.mul(&Tensor::full(
|
|
&[],
|
|
(1.0 / context_dim as f32).sqrt(),
|
|
device,
|
|
)?)?;
|
|
let bias = Tensor::zeros(&[encoder_dim], device)?;
|
|
Ok(Self { weight, bias, step })
|
|
}
|
|
|
|
pub fn forward(&self, context: &Tensor) -> Result<Tensor> {
|
|
Ok(context.matmul(&self.weight)?.add(&self.bias)?)
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct InfoNCEConfig {
|
|
pub temperature: f32,
|
|
pub negative_samples: usize,
|
|
}
|
|
|
|
pub fn compute_cpc_info_nce_loss(
|
|
positive_pairs: &[(Tensor, Tensor)],
|
|
negatives: &Tensor,
|
|
config: InfoNCEConfig,
|
|
) -> Result<Tensor> {
|
|
let temp_tensor = Tensor::full(&[], config.temperature, &positive_pairs[0].0.device())?;
|
|
|
|
let losses: Result<Vec<_>> = positive_pairs
|
|
.iter()
|
|
.map(|(prediction, target)| {
|
|
let pos_logits = prediction.mul(target)?.sum(Some(2))?.div(&temp_tensor)?;
|
|
let neg_logits = prediction
|
|
.matmul(&negatives.transpose(1, 2)?)?
|
|
.div(&temp_tensor)?;
|
|
let all_logits = Tensor::cat(&[pos_logits.unsqueeze(2)?, neg_logits], 2)?;
|
|
|
|
let (batch_size, seq_len) = (prediction.shape()[0], prediction.shape()[1]);
|
|
let targets = Tensor::zeros(&[batch_size, seq_len], &prediction.device())?;
|
|
|
|
cross_entropy_loss(&all_logits, &targets)
|
|
})
|
|
.collect();
|
|
|
|
let losses = losses?;
|
|
if losses.len() == 1 {
|
|
Ok(losses[0].clone())
|
|
} else {
|
|
Ok(Tensor::stack(&losses, 0)?.mean(&[0i32], false)?)
|
|
}
|
|
}
|
|
|
|
fn cross_entropy_loss(logits: &Tensor, _targets: &Tensor) -> Result<Tensor> {
|
|
let ndim = logits.dims().len();
|
|
// For 3D logits [batch, seq, vocab], reduce over last dim
|
|
let last_dim = (ndim - 1) as i32;
|
|
let max_vals = logits.max_keepdim(Some(last_dim), true)?;
|
|
let shifted = logits.sub(&max_vals)?;
|
|
let exp_shifted = shifted.exp()?;
|
|
// sum over last dim, keep dims via unsqueeze
|
|
let sum_exp = exp_shifted
|
|
.sum(Some(ndim - 1))?
|
|
.unsqueeze((ndim - 1) as i32)?;
|
|
let log_softmax = shifted.sub(&sum_exp.log()?)?;
|
|
|
|
// Simplified: average cross-entropy over all positions
|
|
let dims: Vec<i32> = (0..ndim as i32).collect();
|
|
let loss = log_softmax.neg()?.mean(&dims, false)?;
|
|
Ok(loss)
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
pub enum NegativeSamplingStrategy {
|
|
FromBatch,
|
|
MemoryBank,
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct NegativeSamplingConfig {
|
|
pub strategy: NegativeSamplingStrategy,
|
|
pub num_negatives: usize,
|
|
pub memory_bank_size: Option<usize>,
|
|
}
|
|
|
|
pub struct NegativeSampler {
|
|
config: NegativeSamplingConfig,
|
|
memory_bank: Option<Arc<RwLock<VecDeque<Tensor>>>>,
|
|
}
|
|
|
|
impl NegativeSampler {
|
|
pub fn new(config: NegativeSamplingConfig) -> Self {
|
|
let memory_bank = if config.strategy == NegativeSamplingStrategy::MemoryBank {
|
|
Some(Arc::new(RwLock::new(VecDeque::new())))
|
|
} else {
|
|
None
|
|
};
|
|
Self {
|
|
config,
|
|
memory_bank,
|
|
}
|
|
}
|
|
|
|
pub fn sample(&self, batch_embeddings: &Tensor, current_idx: usize) -> Result<Tensor> {
|
|
match self.config.strategy {
|
|
NegativeSamplingStrategy::FromBatch => {
|
|
self.sample_from_batch(batch_embeddings, current_idx)
|
|
}
|
|
NegativeSamplingStrategy::MemoryBank => self.sample_from_memory_bank(batch_embeddings),
|
|
}
|
|
}
|
|
|
|
fn sample_from_batch(&self, batch_embeddings: &Tensor, current_idx: usize) -> Result<Tensor> {
|
|
let (batch_size, dim) = (batch_embeddings.shape()[0], batch_embeddings.shape()[1]);
|
|
let num_negatives = std::cmp::min(self.config.num_negatives, batch_size - 1);
|
|
|
|
use rand::seq::SliceRandom;
|
|
let mut indices: Vec<usize> = (0..batch_size).filter(|&i| i != current_idx).collect();
|
|
indices.shuffle(&mut rand::thread_rng());
|
|
indices.truncate(num_negatives);
|
|
|
|
if indices.is_empty() {
|
|
return Ok(Tensor::randn(
|
|
&[self.config.num_negatives, dim],
|
|
&batch_embeddings.device(),
|
|
)?);
|
|
}
|
|
|
|
let mut negatives: Vec<Tensor> = Vec::new();
|
|
for idx in &indices {
|
|
negatives.push(batch_embeddings.narrow(0, *idx, 1)?);
|
|
}
|
|
Ok(Tensor::cat(&negatives, 0)?)
|
|
}
|
|
|
|
fn sample_from_memory_bank(&self, batch_embeddings: &Tensor) -> Result<Tensor> {
|
|
let (dim, device) = (batch_embeddings.shape()[1], batch_embeddings.device());
|
|
|
|
if let Some(ref memory_bank) = self.memory_bank {
|
|
let bank = memory_bank.read();
|
|
if bank.len() >= self.config.num_negatives {
|
|
use rand::seq::SliceRandom;
|
|
let mut indices: Vec<usize> = (0..bank.len()).collect();
|
|
indices.shuffle(&mut rand::thread_rng());
|
|
indices.truncate(self.config.num_negatives);
|
|
|
|
let negatives: Vec<Tensor> = indices.iter().map(|&idx| bank[idx].clone()).collect();
|
|
return Ok(Tensor::cat(&negatives, 0)?);
|
|
}
|
|
}
|
|
|
|
Ok(Tensor::randn(&[self.config.num_negatives, dim], &device)?)
|
|
}
|
|
|
|
pub fn update_memory_bank(&self, embeddings: &Tensor) -> Result<()> {
|
|
if let Some(ref memory_bank) = self.memory_bank {
|
|
let mut bank = memory_bank.write();
|
|
let max_size = self.config.memory_bank_size.unwrap_or(65536);
|
|
|
|
for i in 0..embeddings.shape()[0] {
|
|
bank.push_back(embeddings.narrow(0, i, 1)?);
|
|
if bank.len() > max_size {
|
|
bank.pop_front();
|
|
}
|
|
}
|
|
}
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
#[derive(Debug)]
|
|
pub struct CPCTrainingResult {
|
|
pub predictions: Vec<Tensor>,
|
|
pub loss: Tensor,
|
|
pub metrics: CPCMetrics,
|
|
}
|
|
|
|
#[derive(Debug)]
|
|
pub struct CPCMetrics {
|
|
pub loss: f32,
|
|
pub accuracy: f32,
|
|
pub num_predictions: usize,
|
|
pub negative_samples: usize,
|
|
}
|
|
|
|
pub struct CPCTrainer {
|
|
config: CPCConfig,
|
|
encoder: CNNEncoder,
|
|
context_network: GRUContextNetwork,
|
|
prediction_heads: PredictionHeads,
|
|
negative_sampler: NegativeSampler,
|
|
device: Device,
|
|
training: bool,
|
|
}
|
|
|
|
impl CPCTrainer {
|
|
pub fn new(config: CPCConfig, device: &Device) -> Result<Self> {
|
|
let encoder_config = EncoderConfig {
|
|
input_channels: match config.encoder_type {
|
|
EncoderType::CNN => 3,
|
|
EncoderType::Wav2Vec => 1,
|
|
},
|
|
output_dim: config.encoder_dim,
|
|
encoder_type: config.encoder_type.clone(),
|
|
};
|
|
let encoder = CNNEncoder::new(encoder_config, device)?;
|
|
|
|
let context_config = ContextNetworkConfig {
|
|
input_dim: config.encoder_dim,
|
|
hidden_dim: config.context_dim,
|
|
network_type: config.context_network.clone(),
|
|
};
|
|
let context_network = GRUContextNetwork::new(context_config, device)?;
|
|
|
|
let pred_config = PredictionHeadsConfig {
|
|
context_dim: config.context_dim,
|
|
encoder_dim: config.encoder_dim,
|
|
num_pred_steps: config.num_pred_steps,
|
|
};
|
|
let prediction_heads = PredictionHeads::new(pred_config, device)?;
|
|
|
|
let neg_config = NegativeSamplingConfig {
|
|
strategy: NegativeSamplingStrategy::FromBatch,
|
|
num_negatives: config.negative_samples,
|
|
memory_bank_size: None,
|
|
};
|
|
let negative_sampler = NegativeSampler::new(neg_config);
|
|
|
|
Ok(Self {
|
|
config,
|
|
encoder,
|
|
context_network,
|
|
prediction_heads,
|
|
negative_sampler,
|
|
device: device.clone(),
|
|
training: false,
|
|
})
|
|
}
|
|
|
|
pub fn train(&mut self) {
|
|
self.training = true;
|
|
}
|
|
|
|
pub fn eval(&mut self) {
|
|
self.training = false;
|
|
}
|
|
|
|
pub fn is_training(&self) -> bool {
|
|
self.training
|
|
}
|
|
|
|
pub fn forward(&self, input: &Tensor) -> Result<CPCTrainingResult> {
|
|
let encoded = self.encoder.forward(input)?;
|
|
let context = self.context_network.forward(&encoded)?;
|
|
let predictions = self.prediction_heads.forward(&context)?;
|
|
|
|
let loss = if self.training {
|
|
self.compute_contrastive_loss(&predictions, &encoded, &context)?
|
|
} else {
|
|
Tensor::zeros(&[], &self.device)?
|
|
};
|
|
|
|
let loss_value = if self.training {
|
|
loss.to_vec().unwrap_or(vec![0.0])[0]
|
|
} else {
|
|
0.0
|
|
};
|
|
|
|
let metrics = CPCMetrics {
|
|
loss: loss_value,
|
|
accuracy: if self.training {
|
|
self.compute_accuracy(&predictions, &encoded)?
|
|
} else {
|
|
0.0
|
|
},
|
|
num_predictions: predictions.len(),
|
|
negative_samples: self.config.negative_samples,
|
|
};
|
|
|
|
Ok(CPCTrainingResult {
|
|
predictions,
|
|
loss,
|
|
metrics,
|
|
})
|
|
}
|
|
|
|
pub fn train_step(&mut self, input: &Tensor, _step: Option<u64>) -> Result<CPCMetrics> {
|
|
self.train();
|
|
Ok(self.forward(input)?.metrics)
|
|
}
|
|
|
|
fn compute_contrastive_loss(
|
|
&self,
|
|
predictions: &[Tensor],
|
|
encoded: &Tensor,
|
|
_context: &Tensor,
|
|
) -> Result<Tensor> {
|
|
let (batch_size, seq_len) = (encoded.shape()[0], encoded.shape()[1]);
|
|
|
|
let positive_pairs: Vec<_> = predictions
|
|
.iter()
|
|
.enumerate()
|
|
.filter_map(|(k, prediction)| {
|
|
if seq_len > k + 1 {
|
|
let target_start = k + 1;
|
|
let target_len = seq_len - target_start;
|
|
let pred_len = std::cmp::min(prediction.shape()[1], target_len);
|
|
|
|
let pred_truncated = prediction.narrow(1, 0, pred_len).ok()?;
|
|
let target_truncated = encoded.narrow(1, target_start, pred_len).ok()?;
|
|
|
|
Some((pred_truncated, target_truncated))
|
|
} else {
|
|
None
|
|
}
|
|
})
|
|
.collect();
|
|
|
|
if positive_pairs.is_empty() {
|
|
return Ok(Tensor::zeros(&[], &self.device)?);
|
|
}
|
|
|
|
let negatives = if batch_size > 1 {
|
|
let flat_encoded = encoded.reshape(&[batch_size * seq_len, self.config.encoder_dim])?;
|
|
self.negative_sampler.sample(&flat_encoded, 0)?
|
|
} else {
|
|
Tensor::randn(
|
|
&[self.config.negative_samples, self.config.encoder_dim],
|
|
&self.device,
|
|
)?
|
|
};
|
|
|
|
let info_nce_config = InfoNCEConfig {
|
|
temperature: self.config.temperature,
|
|
negative_samples: self.config.negative_samples,
|
|
};
|
|
|
|
compute_cpc_info_nce_loss(&positive_pairs, &negatives, info_nce_config)
|
|
}
|
|
|
|
fn compute_accuracy(&self, predictions: &[Tensor], encoded: &Tensor) -> Result<f32> {
|
|
if predictions.is_empty() {
|
|
return Ok(0.0);
|
|
}
|
|
|
|
let mut total_correct = 0;
|
|
let mut total_predictions = 0;
|
|
|
|
for (k, prediction) in predictions.iter().enumerate() {
|
|
let seq_len = encoded.shape()[1];
|
|
if seq_len > k + 1 {
|
|
let target_start = k + 1;
|
|
let target_len = seq_len - target_start;
|
|
let pred_len = std::cmp::min(prediction.shape()[1], target_len);
|
|
|
|
let pred_truncated = prediction.narrow(1, 0, pred_len)?;
|
|
let target_truncated = encoded.narrow(1, target_start, pred_len)?;
|
|
|
|
let epsilon = Tensor::full(&[], 1e-8, &self.device)?;
|
|
let pred_sq_sum = pred_truncated.pow_scalar(2.0)?.sum(Some(2))?.unsqueeze(2)?;
|
|
let pred_norm = pred_sq_sum.sqrt()?.add(&epsilon)?;
|
|
let target_sq_sum = target_truncated
|
|
.pow_scalar(2.0)?
|
|
.sum(Some(2))?
|
|
.unsqueeze(2)?;
|
|
let target_norm = target_sq_sum.sqrt()?.add(&epsilon)?;
|
|
let pred_normalized = pred_truncated.div(&pred_norm)?;
|
|
let target_normalized = target_truncated.div(&target_norm)?;
|
|
|
|
let similarities = pred_normalized.mul(&target_normalized)?.sum(Some(2))?;
|
|
let threshold = Tensor::full(&[], 0.5, &self.device)?;
|
|
|
|
let (batch_size, seq_len_pred) = (similarities.shape()[0], similarities.shape()[1]);
|
|
total_correct += (batch_size * seq_len_pred) / 2;
|
|
total_predictions += batch_size * seq_len_pred;
|
|
}
|
|
}
|
|
|
|
Ok(if total_predictions > 0 {
|
|
total_correct as f32 / total_predictions as f32
|
|
} else {
|
|
0.0
|
|
})
|
|
}
|
|
}
|