//! JEPA (Joint Embedding Predictive Architecture) implementation. //! //! Self-supervised world model for robotics. use embodied_shared::{Observation, WorldModelConfig}; /// JEPA World Model. #[derive(Debug)] pub struct JEPA { /// Observation embedding dimension. embed_dim: usize, /// Hidden dimension. hidden_dim: usize, /// Context encoder weights. context_encoder_weights: Vec, /// Predictor weights. predictor_weights: Vec, /// Target encoder weights (EMA of context encoder). target_encoder_weights: Vec, /// EMA decay rate. ema_decay: f64, /// RNG state. rng_state: u64, } impl JEPA { /// Create a new JEPA model. pub fn new(config: &WorldModelConfig) -> Self { let encoder_size = config.obs_embed_dim * config.hidden_dim; let predictor_size = config.hidden_dim * config.hidden_dim; let mut jepa = Self { embed_dim: config.obs_embed_dim, hidden_dim: config.hidden_dim, context_encoder_weights: vec![0.0; encoder_size], predictor_weights: vec![0.0; predictor_size], target_encoder_weights: vec![0.0; encoder_size], ema_decay: 0.996, rng_state: 42, }; jepa.initialize_weights(); jepa } /// Initialize weights. fn initialize_weights(&mut self) { let scale = (2.0 / self.embed_dim as f64).sqrt(); for w in &mut self.context_encoder_weights { *w = Self::random_normal(&mut self.rng_state) * scale; } for w in &mut self.predictor_weights { *w = Self::random_normal(&mut self.rng_state) * scale * 0.1; } // Initialize target encoder as copy of context encoder self.target_encoder_weights = self.context_encoder_weights.clone(); } /// Random number. fn random(rng: &mut u64) -> f64 { *rng = rng .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); (*rng >> 11) as f64 / (1u64 << 53) as f64 } /// Random normal. fn random_normal(rng: &mut u64) -> f64 { let u1 = Self::random(rng) + 1e-10; let u2 = Self::random(rng); (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos() } /// Encode observation with context encoder. pub fn encode_context(&self, obs: &Observation) -> Vec { self.encode_with_weights(obs, &self.context_encoder_weights) } /// Encode observation with target encoder. pub fn encode_target(&self, obs: &Observation) -> Vec { self.encode_with_weights(obs, &self.target_encoder_weights) } /// Encode observation with given weights. fn encode_with_weights(&self, obs: &Observation, weights: &[f64]) -> Vec { let mut embedding = vec![0.0; self.hidden_dim]; // Simple encoding from joint positions for (i, &pos) in obs.joint_positions.iter().enumerate() { for j in 0..self.hidden_dim { let weight_idx = (i * self.hidden_dim + j) % weights.len(); embedding[j] += pos as f64 * weights[weight_idx]; } } // Apply layer norm (simplified) let mean: f64 = embedding.iter().sum::() / self.hidden_dim as f64; let var: f64 = embedding.iter().map(|x| (x - mean).powi(2)).sum::() / self.hidden_dim as f64; let std = (var + 1e-6).sqrt(); for e in &mut embedding { *e = (*e - mean) / std; } embedding } /// Predict target embedding from context embedding. pub fn predict(&self, context_embedding: &[f64]) -> Vec { let mut prediction = vec![0.0; self.hidden_dim]; for i in 0..self.hidden_dim { for (j, &ctx) in context_embedding.iter().enumerate() { let weight_idx = (i * self.hidden_dim + j) % self.predictor_weights.len(); prediction[i] += ctx * self.predictor_weights[weight_idx]; } prediction[i] = prediction[i].tanh(); } prediction } /// Compute JEPA loss between prediction and target. pub fn compute_loss(&self, prediction: &[f64], target: &[f64]) -> f64 { // MSE loss in embedding space prediction .iter() .zip(target.iter()) .map(|(p, t)| (p - t).powi(2)) .sum::() / prediction.len() as f64 } /// Update target encoder with EMA. pub fn update_target_encoder(&mut self) { for (target, context) in self .target_encoder_weights .iter_mut() .zip(self.context_encoder_weights.iter()) { *target = self.ema_decay * *target + (1.0 - self.ema_decay) * *context; } } /// Train on a pair of observations (context, target). pub fn train_step( &mut self, context_obs: &Observation, target_obs: &Observation, learning_rate: f64, ) -> f64 { // Encode context let context_embedding = self.encode_context(context_obs); // Predict target embedding let prediction = self.predict(&context_embedding); // Encode target (no gradient through target encoder) let target_embedding = self.encode_target(target_obs); // Compute loss let loss = self.compute_loss(&prediction, &target_embedding); // Simplified gradient update for i in 0..self.hidden_dim { let error = prediction[i] - target_embedding[i]; // Update predictor weights for (j, &ctx) in context_embedding.iter().enumerate() { let weight_idx = (i * self.hidden_dim + j) % self.predictor_weights.len(); self.predictor_weights[weight_idx] -= learning_rate * error * ctx * 2.0 / self.hidden_dim as f64; } } // Update context encoder (simplified) for w in &mut self.context_encoder_weights { *w += learning_rate * Self::random_normal(&mut self.rng_state) * 0.0001; } // EMA update of target encoder self.update_target_encoder(); loss } /// Get embedding dimension. #[must_use] pub fn embed_dim(&self) -> usize { self.embed_dim } /// Get hidden dimension. #[must_use] pub fn hidden_dim(&self) -> usize { self.hidden_dim } /// Get number of parameters. #[must_use] pub fn num_parameters(&self) -> usize { self.context_encoder_weights.len() + self.predictor_weights.len() + self.target_encoder_weights.len() } } #[cfg(test)] mod tests { use super::*; use embodied_shared::sample_world_model_config; #[test] fn test_jepa_creation() { let config = sample_world_model_config(); let jepa = JEPA::new(&config); assert_eq!(jepa.hidden_dim, config.hidden_dim); } #[test] fn test_encode_context() { let config = sample_world_model_config(); let jepa = JEPA::new(&config); let obs = embodied_shared::sample_observation(); let embedding = jepa.encode_context(&obs); assert_eq!(embedding.len(), config.hidden_dim); } #[test] fn test_encode_target() { let config = sample_world_model_config(); let jepa = JEPA::new(&config); let obs = embodied_shared::sample_observation(); let embedding = jepa.encode_target(&obs); assert_eq!(embedding.len(), config.hidden_dim); } #[test] fn test_predict() { let config = sample_world_model_config(); let jepa = JEPA::new(&config); let obs = embodied_shared::sample_observation(); let context = jepa.encode_context(&obs); let prediction = jepa.predict(&context); assert_eq!(prediction.len(), config.hidden_dim); } #[test] fn test_compute_loss() { let config = sample_world_model_config(); let jepa = JEPA::new(&config); let prediction = vec![0.5; config.hidden_dim]; let target = vec![0.6; config.hidden_dim]; let loss = jepa.compute_loss(&prediction, &target); assert!(loss > 0.0); assert!(loss < 1.0); } #[test] fn test_train_step() { let config = sample_world_model_config(); let mut jepa = JEPA::new(&config); let obs1 = embodied_shared::sample_observation(); let obs2 = embodied_shared::sample_observation(); let loss = jepa.train_step(&obs1, &obs2, 1e-3); assert!(loss.is_finite()); } #[test] fn test_update_target_encoder() { let config = sample_world_model_config(); let mut jepa = JEPA::new(&config); // Modify context encoder jepa.context_encoder_weights[0] = 1.0; // Update target jepa.update_target_encoder(); // Target should move toward context assert!(jepa.target_encoder_weights[0] > 0.0); } }