Files
rustytorch/demos/rtx-embodied-demo/src/rssm.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
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]>
2026-04-12 07:01:58 -07:00

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);
}
}