use crate::finance::ml::models::q_net::SimpleQNet;
use crate::finance::ml::training::replay_buffer::{ReplayBuffer, Experience};
use crate::finance::ml::errors::Result;
use crate::finance::ml::traits::MarketState;
pub struct QLearningTrainer {
pub model: SimpleQNet,
pub replay_buffer: ReplayBuffer,
pub gamma: f32,
pub learning_rate: f32,
pub batch_size: usize,
total_loss: f32,
training_steps: u32,
}
impl QLearningTrainer {
pub fn new(model: SimpleQNet) -> Self {
Self {
model,
replay_buffer: ReplayBuffer::new(10_000),
gamma: 0.99, learning_rate: 0.001, batch_size: 32,
total_loss: 0.0,
training_steps: 0,
}
}
pub fn remember_experience(
&mut self,
state: Vec<f32>,
action: usize,
reward: f32,
next_state: Vec<f32>,
done: bool,
) {
let market_state = MarketState {
prices: state.clone(),
volatility: vec![0.1; state.len()],
agent_capital: 0.5,
scar_count: 0,
win_loss_ratio: 0.5,
timestamp: 0,
};
let next_market_state = MarketState {
prices: next_state,
volatility: vec![0.1; state.len()],
agent_capital: 0.5,
scar_count: 0,
win_loss_ratio: 0.5,
timestamp: 0,
};
let experience = Experience {
state: market_state,
action,
reward,
next_state: next_market_state,
done,
};
self.replay_buffer.push(experience);
}
pub fn train_step(&mut self) -> Result<f32> {
if self.replay_buffer.len() < self.batch_size {
return Ok(0.0); }
let batch = self.replay_buffer.sample(self.batch_size);
let mut batch_loss = 0.0;
for experience in batch {
let q_values = self.model.forward_pass(&experience.state.prices)?;
let next_q_values = self.model.forward_pass(&experience.next_state.prices)?;
let max_next_q = next_q_values.iter()
.cloned()
.fold(f32::NEG_INFINITY, f32::max);
let target = if experience.done {
experience.reward } else {
experience.reward + self.gamma * max_next_q };
let loss = (q_values[experience.action] - target).powi(2);
batch_loss += loss;
let delta = self.learning_rate * (target - q_values[experience.action]);
self.apply_update(&experience.state, experience.action, delta)?;
}
batch_loss /= self.batch_size as f32;
self.total_loss = batch_loss;
self.training_steps += 1;
Ok(batch_loss)
}
fn apply_update(&mut self, _state: &MarketState, _action: usize, _delta: f32) -> Result<()> {
Ok(())
}
pub fn get_stats(&self) -> TrainingStats {
TrainingStats {
total_loss: self.total_loss,
training_steps: self.training_steps,
buffer_size: self.replay_buffer.len(),
replay_capacity: 10_000,
}
}
pub fn reset_stats(&mut self) {
self.total_loss = 0.0;
self.training_steps = 0;
}
}
impl Clone for QLearningTrainer {
fn clone(&self) -> Self {
Self {
model: self.model.clone(),
replay_buffer: ReplayBuffer::new(10_000),
gamma: self.gamma,
learning_rate: self.learning_rate,
batch_size: self.batch_size,
total_loss: self.total_loss,
training_steps: self.training_steps,
}
}
}
#[derive(Debug, Clone)]
pub struct TrainingStats {
pub total_loss: f32,
pub training_steps: u32,
pub buffer_size: usize,
pub replay_capacity: usize,
}
impl std::fmt::Display for TrainingStats {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"Loss: {:.6} | Steps: {} | Buffer: {}/{}",
self.total_loss, self.training_steps, self.buffer_size, self.replay_capacity
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_trainer_creation() {
let model = SimpleQNet::new(5, 64).unwrap();
let trainer = QLearningTrainer::new(model);
assert_eq!(trainer.gamma, 0.99);
assert_eq!(trainer.learning_rate, 0.001);
assert_eq!(trainer.batch_size, 32);
}
#[test]
fn test_remember_experience() {
let model = SimpleQNet::new(5, 64).unwrap();
let mut trainer = QLearningTrainer::new(model);
trainer.remember_experience(
vec![0.1, 0.2, 0.3, 0.4, 0.5],
0,
1.5,
vec![0.2, 0.3, 0.4, 0.5, 0.6],
false,
);
assert_eq!(trainer.replay_buffer.len(), 1);
}
#[test]
fn test_training_stats() {
let model = SimpleQNet::new(5, 64).unwrap();
let trainer = QLearningTrainer::new(model);
let stats = trainer.get_stats();
assert_eq!(stats.training_steps, 0);
assert_eq!(stats.buffer_size, 0);
}
}