//! RSSM (Recurrent State Space Model) implementation. //! //! Based on DreamerV3 architecture for world modeling. use embodied_shared::{Observation, Transition, WorldModelConfig}; /// RSSM World Model. #[derive(Debug)] pub struct RSSM { /// Observation embedding dimension. obs_embed_dim: usize, /// Action embedding dimension. action_embed_dim: usize, /// Deterministic state dimension. deter_dim: usize, /// Stochastic state dimension. stoch_dim: usize, /// Number of stochastic classes. stoch_classes: usize, /// Hidden dimension. hidden_dim: usize, /// Encoder weights (simplified). encoder_weights: Vec, /// Dynamics weights. dynamics_weights: Vec, /// Decoder weights. decoder_weights: Vec, /// Reward predictor weights. reward_weights: Vec, /// Done predictor weights. done_weights: Vec, /// RNG state. rng_state: u64, } impl RSSM { /// Create a new RSSM model. pub fn new(config: &WorldModelConfig) -> Self { let encoder_size = config.obs_embed_dim * config.hidden_dim; let dynamics_size = (config.deter_dim + config.stoch_dim * config.stoch_classes + config.action_embed_dim) * config.hidden_dim; let decoder_size = (config.deter_dim + config.stoch_dim * config.stoch_classes) * config.hidden_dim; let reward_size = config.hidden_dim * 64; let done_size = config.hidden_dim * 32; let mut rssm = Self { obs_embed_dim: config.obs_embed_dim, action_embed_dim: config.action_embed_dim, deter_dim: config.deter_dim, stoch_dim: config.stoch_dim, stoch_classes: config.stoch_classes, hidden_dim: config.hidden_dim, encoder_weights: vec![0.0; encoder_size], dynamics_weights: vec![0.0; dynamics_size], decoder_weights: vec![0.0; decoder_size], reward_weights: vec![0.0; reward_size], done_weights: vec![0.0; done_size], rng_state: 42, }; rssm.initialize_weights(); rssm } /// Initialize weights with Xavier initialization. fn initialize_weights(&mut self) { let init_weights = |weights: &mut [f64], rng: &mut u64, fan_in: usize| { let scale = (2.0 / fan_in as f64).sqrt(); for w in weights.iter_mut() { *w = Self::random_normal(rng) * scale; } }; init_weights( &mut self.encoder_weights, &mut self.rng_state, self.obs_embed_dim, ); init_weights( &mut self.dynamics_weights, &mut self.rng_state, self.deter_dim, ); init_weights( &mut self.decoder_weights, &mut self.rng_state, self.deter_dim, ); init_weights( &mut self.reward_weights, &mut self.rng_state, self.hidden_dim, ); init_weights(&mut self.done_weights, &mut self.rng_state, self.hidden_dim); } /// Random number generator. fn random(rng: &mut u64) -> f64 { *rng = rng .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); (*rng >> 11) as f64 / (1u64 << 53) as f64 } /// Random normal using Box-Muller. 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 to state. pub fn encode_observation(&self, obs: &Observation) -> Vec { let state_dim = self.deter_dim + self.stoch_dim * self.stoch_classes; let mut state = vec![0.0; state_dim]; // Simple encoding: use joint positions and velocities for (i, &pos) in obs.joint_positions.iter().enumerate() { let idx = i % state_dim; state[idx] += pos as f64 * 0.1; } for (i, &vel) in obs.joint_velocities.iter().enumerate() { let idx = (i + obs.joint_positions.len()) % state_dim; state[idx] += vel as f64 * 0.01; } // Apply encoder weights for (i, s) in state.iter_mut().enumerate() { let weight_idx = i % self.encoder_weights.len(); *s = (*s * self.encoder_weights[weight_idx]).tanh(); } state } /// Imagine one step forward. pub fn imagine_step(&self, state: &[f64], action: &[f64]) -> (Vec, f64, f64) { let state_dim = state.len(); let mut next_state = vec![0.0; state_dim]; // GRU-style update for i in 0..state_dim { let action_idx = i % action.len(); let weight_idx = i % self.dynamics_weights.len(); // Reset gate let reset = (state[i] * self.dynamics_weights[weight_idx] + action[action_idx] * 0.1).tanh(); // Update gate let update = ((state[i] + action[action_idx] * 0.1) * self.dynamics_weights[(weight_idx + 1) % self.dynamics_weights.len()]) .sigmoid(); // Candidate let candidate = (reset * state[i] + action[action_idx] * 0.2 + self.dynamics_weights[(weight_idx + 2) % self.dynamics_weights.len()]) .tanh(); next_state[i] = (1.0 - update) * state[i] + update * candidate; } // Predict reward let reward = self.predict_reward(&next_state); // Predict done probability let done_logit = self.predict_done(&next_state); let done_prob = done_logit.sigmoid(); (next_state, reward, done_prob) } /// Predict reward from state. fn predict_reward(&self, state: &[f64]) -> f64 { let mut reward = 0.0; for (i, &s) in state.iter().enumerate() { let weight_idx = i % self.reward_weights.len(); reward += s * self.reward_weights[weight_idx]; } reward.tanh() // Bounded reward } /// Predict done from state. fn predict_done(&self, state: &[f64]) -> f64 { let mut logit = 0.0; for (i, &s) in state.iter().enumerate() { let weight_idx = i % self.done_weights.len(); logit += s * self.done_weights[weight_idx]; } logit * 0.01 // Small logit, mostly not done } /// Train on a batch of sequences. pub fn train_step(&mut self, batch: &[Vec], learning_rate: f64) -> f64 { let mut total_loss = 0.0; for sequence in batch { let loss = self.train_sequence(sequence, learning_rate); total_loss += loss; } total_loss / batch.len() as f64 } /// Train on a single sequence. fn train_sequence(&mut self, sequence: &[Transition], learning_rate: f64) -> f64 { let mut loss = 0.0; for trans in sequence { // Encode current observation let state = self.encode_observation(&trans.observation); // Encode action let action: Vec = trans.action.joint_positions.as_ref().map_or_else( || vec![0.0; self.action_embed_dim], |p| p.iter().map(|&x| x as f64).collect(), ); // Imagine next state let (pred_next, pred_reward, pred_done) = self.imagine_step(&state, &action); // Encode actual next observation let actual_next = self.encode_observation(&trans.next_observation); // Compute losses let state_loss: f64 = pred_next .iter() .zip(actual_next.iter()) .map(|(p, a)| (p - a).powi(2)) .sum::() / pred_next.len() as f64; let reward_loss = (pred_reward - trans.reward as f64).powi(2); let done_target = if trans.done { 1.0 } else { 0.0 }; let done_loss = -done_target * pred_done.ln().max(-10.0) - (1.0 - done_target) * (1.0 - pred_done).ln().max(-10.0); loss += state_loss + reward_loss * 0.5 + done_loss * 0.1; // Simplified gradient update (in real implementation, use autograd) self.update_weights(learning_rate * 0.01); } loss / sequence.len() as f64 } /// Simplified weight update. fn update_weights(&mut self, lr: f64) { // Add small random perturbation as simplified training for w in &mut self.encoder_weights { *w += lr * Self::random_normal(&mut self.rng_state) * 0.001; } for w in &mut self.dynamics_weights { *w += lr * Self::random_normal(&mut self.rng_state) * 0.001; } } /// Get state dimension. #[must_use] pub fn state_dim(&self) -> usize { self.deter_dim + self.stoch_dim * self.stoch_classes } /// Get number of parameters. #[must_use] pub fn num_parameters(&self) -> usize { self.encoder_weights.len() + self.dynamics_weights.len() + self.decoder_weights.len() + self.reward_weights.len() + self.done_weights.len() } } /// Sigmoid activation. trait Sigmoid { fn sigmoid(self) -> Self; } impl Sigmoid for f64 { fn sigmoid(self) -> Self { 1.0 / (1.0 + (-self).exp()) } } #[cfg(test)] mod tests { use super::*; use embodied_shared::sample_world_model_config; #[test] fn test_rssm_creation() { let config = sample_world_model_config(); let rssm = RSSM::new(&config); assert_eq!(rssm.deter_dim, config.deter_dim); assert_eq!(rssm.stoch_dim, config.stoch_dim); } #[test] fn test_encode_observation() { let config = sample_world_model_config(); let rssm = RSSM::new(&config); let obs = embodied_shared::sample_observation(); let state = rssm.encode_observation(&obs); let expected_dim = config.deter_dim + config.stoch_dim * config.stoch_classes; assert_eq!(state.len(), expected_dim); } #[test] fn test_imagine_step() { let config = sample_world_model_config(); let rssm = RSSM::new(&config); let obs = embodied_shared::sample_observation(); let state = rssm.encode_observation(&obs); let action = vec![0.1; 7]; let (next_state, reward, done) = rssm.imagine_step(&state, &action); assert_eq!(next_state.len(), state.len()); assert!(reward.is_finite()); assert!(done >= 0.0 && done <= 1.0); } #[test] fn test_train_sequence() { let config = sample_world_model_config(); let mut rssm = RSSM::new(&config); let sequence = vec![ embodied_shared::sample_transition(), embodied_shared::sample_transition(), ]; let loss = rssm.train_sequence(&sequence, 1e-4); assert!(loss.is_finite()); } #[test] fn test_state_dim() { let config = sample_world_model_config(); let rssm = RSSM::new(&config); let expected = config.deter_dim + config.stoch_dim * config.stoch_classes; assert_eq!(rssm.state_dim(), expected); } #[test] fn test_num_parameters() { let config = sample_world_model_config(); let rssm = RSSM::new(&config); assert!(rssm.num_parameters() > 0); } }