//! 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::() / 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); }