351 lines
9.2 KiB
Rust
351 lines
9.2 KiB
Rust
//! Tests for RLHF (Reinforcement Learning from Human Feedback) components
|
|
|
|
use rtx_rl::rlhf::{
|
|
PPOConfig, PPOTrainer, PreferenceDataset, PreferencePair, RLHFConfig, RLHFTrainer, RewardModel,
|
|
RewardModelConfig,
|
|
};
|
|
|
|
#[test]
|
|
fn test_reward_model_creation() {
|
|
let config = RewardModelConfig::new(768, 256, 4);
|
|
let model = RewardModel::new(config);
|
|
|
|
assert_eq!(model.input_dim(), 768);
|
|
assert_eq!(model.hidden_dim(), 256);
|
|
assert_eq!(model.num_layers(), 4);
|
|
}
|
|
|
|
#[test]
|
|
fn test_reward_model_forward() {
|
|
let config = RewardModelConfig::new(64, 32, 2);
|
|
let model = RewardModel::new(config);
|
|
|
|
let input = vec![0.5; 64];
|
|
let reward = model.forward(&input);
|
|
|
|
assert!(reward.is_finite());
|
|
assert!(reward >= -10.0 && reward <= 10.0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_reward_model_batch_forward() {
|
|
let config = RewardModelConfig::new(64, 32, 2);
|
|
let model = RewardModel::new(config);
|
|
|
|
let batch = vec![vec![0.1; 64], vec![0.5; 64], vec![0.9; 64]];
|
|
|
|
let rewards = model.forward_batch(&batch);
|
|
assert_eq!(rewards.len(), 3);
|
|
|
|
for reward in &rewards {
|
|
assert!(reward.is_finite());
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_reward_model_training() {
|
|
let config = RewardModelConfig::new(64, 32, 2).with_learning_rate(0.001);
|
|
let mut model = RewardModel::new(config);
|
|
|
|
// Create preference pairs
|
|
let pairs = vec![
|
|
PreferencePair {
|
|
chosen: vec![0.8; 64],
|
|
rejected: vec![0.2; 64],
|
|
},
|
|
PreferencePair {
|
|
chosen: vec![0.9; 64],
|
|
rejected: vec![0.1; 64],
|
|
},
|
|
];
|
|
|
|
let initial_loss = model.compute_loss(&pairs);
|
|
model.train_step(&pairs);
|
|
let final_loss = model.compute_loss(&pairs);
|
|
|
|
assert!(final_loss < initial_loss);
|
|
}
|
|
|
|
#[test]
|
|
fn test_ppo_config() {
|
|
let config = PPOConfig::new(64, 32, 2)
|
|
.with_clip_epsilon(0.2)
|
|
.with_value_coef(0.5)
|
|
.with_entropy_coef(0.01)
|
|
.with_gae_lambda(0.95)
|
|
.with_discount_factor(0.99);
|
|
|
|
assert_eq!(config.state_dim(), 64);
|
|
assert_eq!(config.action_dim(), 32);
|
|
assert_eq!(config.hidden_dim(), 2);
|
|
assert_eq!(config.clip_epsilon(), 0.2);
|
|
assert_eq!(config.value_coef(), 0.5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_ppo_trainer_creation() {
|
|
let config = PPOConfig::new(64, 32, 128);
|
|
let trainer = PPOTrainer::new(config);
|
|
|
|
assert_eq!(trainer.num_epochs(), 4);
|
|
assert_eq!(trainer.batch_size(), 64);
|
|
}
|
|
|
|
#[test]
|
|
fn test_ppo_compute_advantages() {
|
|
let config = PPOConfig::new(64, 32, 128);
|
|
let trainer = PPOTrainer::new(config);
|
|
|
|
let rewards = vec![1.0, 0.5, 2.0, 1.5];
|
|
let values = vec![0.8, 0.6, 1.8, 1.4];
|
|
let next_value = 1.2;
|
|
|
|
let advantages = trainer.compute_advantages(&rewards, &values, next_value);
|
|
|
|
assert_eq!(advantages.len(), 4);
|
|
for adv in &advantages {
|
|
assert!(adv.is_finite());
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_ppo_compute_returns() {
|
|
let config = PPOConfig::new(64, 32, 128).with_discount_factor(0.99);
|
|
let trainer = PPOTrainer::new(config);
|
|
|
|
let rewards = vec![1.0, 0.5, 2.0];
|
|
let returns = trainer.compute_returns(&rewards);
|
|
|
|
assert_eq!(returns.len(), 3);
|
|
assert!(returns[0] > returns[1]); // First return should include future rewards
|
|
}
|
|
|
|
#[test]
|
|
fn test_ppo_training_step() {
|
|
let config = PPOConfig::new(4, 2, 32);
|
|
let mut trainer = PPOTrainer::new(config);
|
|
|
|
let states = vec![vec![0.1, 0.2, 0.3, 0.4], vec![0.5, 0.6, 0.7, 0.8]];
|
|
let actions = vec![0, 1];
|
|
let rewards = vec![1.0, 0.5];
|
|
let old_log_probs = vec![-0.5, -0.7];
|
|
let advantages = vec![0.5, -0.3];
|
|
let returns = vec![1.5, 0.5];
|
|
|
|
let metrics = trainer.train_step(
|
|
&states,
|
|
&actions,
|
|
&rewards,
|
|
&old_log_probs,
|
|
&advantages,
|
|
&returns,
|
|
);
|
|
|
|
assert!(metrics.policy_loss.is_finite());
|
|
assert!(metrics.value_loss.is_finite());
|
|
assert!(metrics.entropy.is_finite());
|
|
assert!(metrics.kl_divergence >= 0.0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_preference_dataset_creation() {
|
|
let dataset = PreferenceDataset::new();
|
|
assert_eq!(dataset.size(), 0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_preference_dataset_add_pair() {
|
|
let mut dataset = PreferenceDataset::new();
|
|
|
|
let pair = PreferencePair {
|
|
chosen: vec![0.8; 64],
|
|
rejected: vec![0.2; 64],
|
|
};
|
|
|
|
dataset.add_pair(pair.clone());
|
|
assert_eq!(dataset.size(), 1);
|
|
|
|
let retrieved = dataset.get_batch(1);
|
|
assert_eq!(retrieved.len(), 1);
|
|
assert_eq!(retrieved[0].chosen, pair.chosen);
|
|
}
|
|
|
|
#[test]
|
|
fn test_preference_dataset_sampling() {
|
|
let mut dataset = PreferenceDataset::new();
|
|
|
|
for i in 0..10 {
|
|
let pair = PreferencePair {
|
|
chosen: vec![i as f32 * 0.1; 64],
|
|
rejected: vec![i as f32 * 0.05; 64],
|
|
};
|
|
dataset.add_pair(pair);
|
|
}
|
|
|
|
let batch = dataset.sample_batch(5);
|
|
assert_eq!(batch.len(), 5);
|
|
|
|
// Check that samples are different
|
|
let first_sum: f32 = batch[0].chosen.iter().sum();
|
|
let mut all_same = true;
|
|
for pair in &batch[1..] {
|
|
let sum: f32 = pair.chosen.iter().sum();
|
|
if (sum - first_sum).abs() > 0.001 {
|
|
all_same = false;
|
|
break;
|
|
}
|
|
}
|
|
assert!(!all_same);
|
|
}
|
|
|
|
#[test]
|
|
fn test_preference_dataset_clear() {
|
|
let mut dataset = PreferenceDataset::new();
|
|
|
|
for i in 0..5 {
|
|
let pair = PreferencePair {
|
|
chosen: vec![i as f32; 64],
|
|
rejected: vec![i as f32 * 0.5; 64],
|
|
};
|
|
dataset.add_pair(pair);
|
|
}
|
|
|
|
assert_eq!(dataset.size(), 5);
|
|
dataset.clear();
|
|
assert_eq!(dataset.size(), 0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_rlhf_trainer_creation() {
|
|
let config = RLHFConfig::new(64, 32, 128)
|
|
.with_reward_learning_rate(0.0001)
|
|
.with_policy_learning_rate(0.0003)
|
|
.with_num_reward_epochs(1)
|
|
.with_num_ppo_epochs(4);
|
|
|
|
let trainer = RLHFTrainer::new(config);
|
|
assert_eq!(trainer.num_reward_epochs(), 1);
|
|
assert_eq!(trainer.num_ppo_epochs(), 4);
|
|
}
|
|
|
|
#[test]
|
|
fn test_rlhf_train_reward_model() {
|
|
let config = RLHFConfig::new(64, 32, 128);
|
|
let mut trainer = RLHFTrainer::new(config);
|
|
|
|
let pairs = vec![
|
|
PreferencePair {
|
|
chosen: vec![0.9; 64],
|
|
rejected: vec![0.1; 64],
|
|
},
|
|
PreferencePair {
|
|
chosen: vec![0.8; 64],
|
|
rejected: vec![0.2; 64],
|
|
},
|
|
];
|
|
|
|
let metrics = trainer.train_reward_model(&pairs);
|
|
assert!(metrics.loss > 0.0);
|
|
assert!(metrics.accuracy >= 0.0 && metrics.accuracy <= 1.0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_rlhf_generate_rewards() {
|
|
let config = RLHFConfig::new(64, 32, 128);
|
|
let trainer = RLHFTrainer::new(config);
|
|
|
|
let states = vec![vec![0.1; 64], vec![0.5; 64], vec![0.9; 64]];
|
|
|
|
let rewards = trainer.generate_rewards(&states);
|
|
assert_eq!(rewards.len(), 3);
|
|
|
|
for reward in &rewards {
|
|
assert!(reward.is_finite());
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_rlhf_full_training_loop() {
|
|
let config = RLHFConfig::new(4, 2, 32)
|
|
.with_num_reward_epochs(1)
|
|
.with_num_ppo_epochs(1);
|
|
let mut trainer = RLHFTrainer::new(config);
|
|
|
|
// Step 1: Train reward model with preferences
|
|
let preferences = vec![PreferencePair {
|
|
chosen: vec![0.9, 0.8, 0.7, 0.6],
|
|
rejected: vec![0.1, 0.2, 0.3, 0.4],
|
|
}];
|
|
|
|
let reward_metrics = trainer.train_reward_model(&preferences);
|
|
assert!(reward_metrics.loss > 0.0);
|
|
|
|
// Step 2: Generate rewards for new states
|
|
let states = vec![vec![0.5, 0.5, 0.5, 0.5], vec![0.7, 0.7, 0.7, 0.7]];
|
|
|
|
let rewards = trainer.generate_rewards(&states);
|
|
assert_eq!(rewards.len(), 2);
|
|
|
|
// Step 3: Train policy with PPO
|
|
let actions = vec![0, 1];
|
|
let old_log_probs = vec![-0.5, -0.6];
|
|
|
|
let ppo_metrics = trainer.train_policy(&states, &actions, &rewards, &old_log_probs);
|
|
|
|
assert!(ppo_metrics.policy_loss.is_finite());
|
|
assert!(ppo_metrics.value_loss.is_finite());
|
|
}
|
|
|
|
#[test]
|
|
fn test_rlhf_adaptive_training() {
|
|
let config = RLHFConfig::new(4, 2, 32)
|
|
.with_adaptive_kl_target(0.01)
|
|
.with_kl_penalty_coef(0.1);
|
|
|
|
let mut trainer = RLHFTrainer::new(config);
|
|
|
|
// Simulate high KL divergence
|
|
trainer.update_kl_penalty(0.05);
|
|
assert!(trainer.kl_penalty_coef() > 0.1);
|
|
|
|
// Simulate low KL divergence
|
|
trainer.update_kl_penalty(0.005);
|
|
assert!(trainer.kl_penalty_coef() < 0.15);
|
|
}
|
|
|
|
#[test]
|
|
fn test_reward_model_normalization() {
|
|
let config = RewardModelConfig::new(64, 32, 2).with_normalize_rewards(true);
|
|
let model = RewardModel::new(config);
|
|
|
|
let batch = vec![vec![0.1; 64], vec![0.5; 64], vec![0.9; 64]];
|
|
|
|
let rewards = model.forward_batch(&batch);
|
|
let mean: f32 = rewards.iter().sum::<f32>() / rewards.len() as f32;
|
|
|
|
// Normalized rewards should have mean close to 0
|
|
assert!(mean.abs() < 1.0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_ppo_early_stopping() {
|
|
let config = PPOConfig::new(4, 2, 32).with_early_stop_kl(0.02);
|
|
let mut trainer = PPOTrainer::new(config);
|
|
|
|
let _states = vec![vec![0.5; 4]; 10];
|
|
let _actions = vec![0; 10];
|
|
let _rewards = vec![1.0; 10];
|
|
let _old_log_probs = vec![-0.5; 10];
|
|
let _advantages = vec![0.5; 10];
|
|
let _returns = vec![1.5; 10];
|
|
|
|
// Simulate training with increasing KL
|
|
trainer.set_kl_divergence(0.01);
|
|
let should_continue = trainer.should_continue_training();
|
|
assert!(should_continue);
|
|
|
|
trainer.set_kl_divergence(0.03);
|
|
let should_stop = !trainer.should_continue_training();
|
|
assert!(should_stop);
|
|
}
|