Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
+296
View File
@@ -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);
}
}