Initial commit
This commit is contained in:
@@ -0,0 +1,197 @@
|
||||
//! Reward model for RLHF
|
||||
|
||||
use crate::rlhf::preference_dataset::PreferencePair;
|
||||
use std::sync::Arc;
|
||||
use std::sync::Mutex;
|
||||
|
||||
/// Configuration for reward model
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RewardModelConfig {
|
||||
input_dim: usize,
|
||||
hidden_dim: usize,
|
||||
num_layers: usize,
|
||||
learning_rate: f32,
|
||||
normalize_rewards: bool,
|
||||
}
|
||||
|
||||
impl RewardModelConfig {
|
||||
pub fn new(input_dim: usize, hidden_dim: usize, num_layers: usize) -> Self {
|
||||
Self {
|
||||
input_dim,
|
||||
hidden_dim,
|
||||
num_layers,
|
||||
learning_rate: 0.0001,
|
||||
normalize_rewards: false,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_learning_rate(mut self, lr: f32) -> Self {
|
||||
self.learning_rate = lr;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_normalize_rewards(mut self, normalize: bool) -> Self {
|
||||
self.normalize_rewards = normalize;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Reward model for scoring trajectories
|
||||
#[derive(Debug)]
|
||||
pub struct RewardModel {
|
||||
config: RewardModelConfig,
|
||||
weights: Arc<Mutex<Vec<Vec<Vec<f32>>>>>,
|
||||
biases: Arc<Mutex<Vec<Vec<f32>>>>,
|
||||
}
|
||||
|
||||
impl RewardModel {
|
||||
pub fn new(config: RewardModelConfig) -> Self {
|
||||
let mut weights = Vec::new();
|
||||
let mut biases = Vec::new();
|
||||
|
||||
// Initialize network layers
|
||||
let mut prev_dim = config.input_dim;
|
||||
for _ in 0..config.num_layers {
|
||||
// Initialize weights with small random values
|
||||
let layer_weights: Vec<Vec<f32>> = (0..config.hidden_dim)
|
||||
.map(|_| {
|
||||
(0..prev_dim)
|
||||
.map(|_| (rand::random::<f32>() - 0.5) * 0.1)
|
||||
.collect()
|
||||
})
|
||||
.collect();
|
||||
weights.push(layer_weights);
|
||||
|
||||
// Initialize biases to zero
|
||||
biases.push(vec![0.0; config.hidden_dim]);
|
||||
prev_dim = config.hidden_dim;
|
||||
}
|
||||
|
||||
// Output layer (single reward value)
|
||||
let output_weights: Vec<Vec<f32>> = vec![
|
||||
(0..prev_dim)
|
||||
.map(|_| (rand::random::<f32>() - 0.5) * 0.1)
|
||||
.collect(),
|
||||
];
|
||||
weights.push(output_weights);
|
||||
biases.push(vec![0.0]);
|
||||
|
||||
Self {
|
||||
config,
|
||||
weights: Arc::new(Mutex::new(weights)),
|
||||
biases: Arc::new(Mutex::new(biases)),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn input_dim(&self) -> usize {
|
||||
self.config.input_dim
|
||||
}
|
||||
|
||||
pub fn hidden_dim(&self) -> usize {
|
||||
self.config.hidden_dim
|
||||
}
|
||||
|
||||
pub fn num_layers(&self) -> usize {
|
||||
self.config.num_layers
|
||||
}
|
||||
|
||||
pub fn forward(&self, input: &[f32]) -> f32 {
|
||||
let weights = self.weights.lock().unwrap();
|
||||
let biases = self.biases.lock().unwrap();
|
||||
|
||||
let mut activation = input.to_vec();
|
||||
|
||||
// Forward through hidden layers
|
||||
for layer_idx in 0..self.config.num_layers {
|
||||
let mut next_activation = vec![0.0; self.config.hidden_dim];
|
||||
|
||||
for (i, neuron_weights) in weights[layer_idx].iter().enumerate() {
|
||||
let mut sum = biases[layer_idx][i];
|
||||
for (j, &w) in neuron_weights.iter().enumerate() {
|
||||
sum += w * activation[j];
|
||||
}
|
||||
// ReLU activation
|
||||
next_activation[i] = sum.max(0.0);
|
||||
}
|
||||
|
||||
activation = next_activation;
|
||||
}
|
||||
|
||||
// Output layer (no activation, raw reward value)
|
||||
let output_idx = self.config.num_layers;
|
||||
let mut reward = biases[output_idx][0];
|
||||
for (i, &w) in weights[output_idx][0].iter().enumerate() {
|
||||
reward += w * activation[i];
|
||||
}
|
||||
|
||||
// Clip reward to reasonable range
|
||||
reward.max(-10.0).min(10.0)
|
||||
}
|
||||
|
||||
pub fn forward_batch(&self, batch: &[Vec<f32>]) -> Vec<f32> {
|
||||
let mut rewards: Vec<f32> = batch.iter().map(|input| self.forward(input)).collect();
|
||||
|
||||
if self.config.normalize_rewards && rewards.len() > 1 {
|
||||
// Normalize rewards to have mean 0
|
||||
let mean = rewards.iter().sum::<f32>() / rewards.len() as f32;
|
||||
let std = {
|
||||
let variance =
|
||||
rewards.iter().map(|&r| (r - mean).powi(2)).sum::<f32>() / rewards.len() as f32;
|
||||
variance.sqrt().max(1e-8)
|
||||
};
|
||||
|
||||
for reward in &mut rewards {
|
||||
*reward = (*reward - mean) / std;
|
||||
}
|
||||
}
|
||||
|
||||
rewards
|
||||
}
|
||||
|
||||
pub fn compute_loss(&self, pairs: &[PreferencePair]) -> f32 {
|
||||
let mut total_loss = 0.0;
|
||||
|
||||
for pair in pairs {
|
||||
let chosen_reward = self.forward(&pair.chosen);
|
||||
let rejected_reward = self.forward(&pair.rejected);
|
||||
|
||||
// Bradley-Terry model loss
|
||||
let diff = chosen_reward - rejected_reward;
|
||||
let loss = -(diff / (1.0 + (-diff).exp()).ln());
|
||||
total_loss += loss;
|
||||
}
|
||||
|
||||
total_loss / pairs.len() as f32
|
||||
}
|
||||
|
||||
pub fn train_step(&mut self, pairs: &[PreferencePair]) {
|
||||
// Simple gradient descent update
|
||||
let mut weights = self.weights.lock().unwrap();
|
||||
let mut biases = self.biases.lock().unwrap();
|
||||
|
||||
// Compute gradients (simplified)
|
||||
for pair in pairs {
|
||||
let chosen_reward = self.forward(&pair.chosen);
|
||||
let rejected_reward = self.forward(&pair.rejected);
|
||||
|
||||
let diff = chosen_reward - rejected_reward;
|
||||
let grad = 1.0 / (1.0 + diff.exp());
|
||||
|
||||
// Update weights (simplified gradient update)
|
||||
for layer_weights in weights.iter_mut() {
|
||||
for neuron_weights in layer_weights.iter_mut() {
|
||||
for weight in neuron_weights.iter_mut() {
|
||||
*weight += self.config.learning_rate * grad * 0.01;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Update biases
|
||||
for layer_biases in biases.iter_mut() {
|
||||
for bias in layer_biases.iter_mut() {
|
||||
*bias += self.config.learning_rate * grad * 0.01;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user