Consistent formatting pass: line wrapping, import sorting, trailing whitespace removal, let-chain indentation, merged derive attributes, and unsafe block reformatting. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
366 lines
11 KiB
Rust
366 lines
11 KiB
Rust
//! 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<f64>,
|
|
/// Dynamics weights.
|
|
dynamics_weights: Vec<f64>,
|
|
/// Decoder weights.
|
|
decoder_weights: Vec<f64>,
|
|
/// Reward predictor weights.
|
|
reward_weights: Vec<f64>,
|
|
/// Done predictor weights.
|
|
done_weights: Vec<f64>,
|
|
/// 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<f64> {
|
|
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, 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<Transition>], 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<f64> = 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::<f64>()
|
|
/ 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);
|
|
}
|
|
}
|