Initial commit
This commit is contained in:
@@ -0,0 +1,296 @@
|
||||
//! 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<f64>,
|
||||
/// Predictor weights.
|
||||
predictor_weights: Vec<f64>,
|
||||
/// Target encoder weights (EMA of context encoder).
|
||||
target_encoder_weights: Vec<f64>,
|
||||
/// 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<f64> {
|
||||
self.encode_with_weights(obs, &self.context_encoder_weights)
|
||||
}
|
||||
|
||||
/// Encode observation with target encoder.
|
||||
pub fn encode_target(&self, obs: &Observation) -> Vec<f64> {
|
||||
self.encode_with_weights(obs, &self.target_encoder_weights)
|
||||
}
|
||||
|
||||
/// Encode observation with given weights.
|
||||
fn encode_with_weights(&self, obs: &Observation, weights: &[f64]) -> Vec<f64> {
|
||||
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::<f64>() / self.hidden_dim as f64;
|
||||
let var: f64 =
|
||||
embedding.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / 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<f64> {
|
||||
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::<f64>()
|
||||
/ 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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user