use crate::error::Result;
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use std::collections::VecDeque;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RLMarketState {
pub inventory: Decimal,
pub mid_price: Decimal,
pub volatility: Decimal,
pub spread: Decimal,
pub order_book_imbalance: Decimal,
pub recent_pnl: Decimal,
pub time_elapsed: Decimal,
pub trend: Decimal,
}
impl RLMarketState {
pub fn normalize(&self, max_inventory: Decimal, max_price: Decimal) -> Vec<f64> {
vec![
(self.inventory / max_inventory)
.to_string()
.parse()
.unwrap_or(0.0),
(self.mid_price / max_price)
.to_string()
.parse()
.unwrap_or(0.0),
self.volatility.to_string().parse().unwrap_or(0.0),
(self.spread / max_price).to_string().parse().unwrap_or(0.0),
self.order_book_imbalance.to_string().parse().unwrap_or(0.0),
(self.recent_pnl / max_price)
.to_string()
.parse()
.unwrap_or(0.0),
self.time_elapsed.to_string().parse().unwrap_or(0.0),
self.trend.to_string().parse().unwrap_or(0.0),
]
}
pub fn dimension() -> usize {
8
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MarketAction {
pub bid_offset: Decimal,
pub ask_offset: Decimal,
pub bid_quantity: Decimal,
pub ask_quantity: Decimal,
}
impl MarketAction {
pub fn from_continuous(values: &[f64]) -> Self {
let bid_offset_raw = values.first().copied().unwrap_or(0.0);
let ask_offset_raw = values.get(1).copied().unwrap_or(0.0);
let bid_offset = Decimal::from_f64_retain((1.0 / (1.0 + (-bid_offset_raw).exp())) * 0.02)
.unwrap_or(Decimal::from_f64_retain(0.005).unwrap());
let ask_offset = Decimal::from_f64_retain((1.0 / (1.0 + (-ask_offset_raw).exp())) * 0.02)
.unwrap_or(Decimal::from_f64_retain(0.005).unwrap());
let bid_qty_raw = values.get(2).copied().unwrap_or(0.0);
let ask_qty_raw = values.get(3).copied().unwrap_or(0.0);
let bid_quantity =
Decimal::from_f64_retain(50.0 + (1.0 / (1.0 + (-bid_qty_raw).exp())) * 150.0)
.unwrap_or(Decimal::from(100));
let ask_quantity =
Decimal::from_f64_retain(50.0 + (1.0 / (1.0 + (-ask_qty_raw).exp())) * 150.0)
.unwrap_or(Decimal::from(100));
Self {
bid_offset,
ask_offset,
bid_quantity,
ask_quantity,
}
}
pub fn dimension() -> usize {
4
}
}
#[derive(Debug, Clone)]
pub struct Experience {
pub state: RLMarketState,
pub action: MarketAction,
pub reward: Decimal,
pub next_state: RLMarketState,
pub done: bool,
}
pub struct ReplayBuffer {
buffer: VecDeque<Experience>,
capacity: usize,
}
impl ReplayBuffer {
pub fn new(capacity: usize) -> Self {
Self {
buffer: VecDeque::with_capacity(capacity),
capacity,
}
}
pub fn push(&mut self, experience: Experience) {
if self.buffer.len() >= self.capacity {
self.buffer.pop_front();
}
self.buffer.push_back(experience);
}
pub fn sample(&self, batch_size: usize) -> Vec<Experience> {
use rand::RngExt;
let mut rng = rand::rng();
let buffer_vec: Vec<_> = self.buffer.iter().collect();
let sample_size = batch_size.min(self.buffer.len());
let mut sampled = Vec::with_capacity(sample_size);
let mut indices: Vec<usize> = (0..buffer_vec.len()).collect();
for _ in 0..sample_size {
if indices.is_empty() {
break;
}
let idx = rng.random_range(0..indices.len());
let buffer_idx = indices.swap_remove(idx);
sampled.push(buffer_vec[buffer_idx].clone());
}
sampled
}
pub fn len(&self) -> usize {
self.buffer.len()
}
pub fn is_empty(&self) -> bool {
self.buffer.is_empty()
}
}
#[derive(Debug, Clone)]
pub struct RLMarketMakerConfig {
pub learning_rate: f64,
pub discount_factor: f64,
pub epsilon: f64,
pub epsilon_decay: f64,
pub epsilon_min: f64,
pub target_update_freq: usize,
pub batch_size: usize,
pub replay_buffer_capacity: usize,
pub inventory_penalty: Decimal,
pub adverse_selection_penalty: Decimal,
}
impl Default for RLMarketMakerConfig {
fn default() -> Self {
Self {
learning_rate: 0.001,
discount_factor: 0.99,
epsilon: 1.0,
epsilon_decay: 0.995,
epsilon_min: 0.01,
target_update_freq: 100,
batch_size: 32,
replay_buffer_capacity: 10000,
inventory_penalty: Decimal::from_f64_retain(0.01).unwrap(),
adverse_selection_penalty: Decimal::from_f64_retain(0.005).unwrap(),
}
}
}
#[derive(Debug, Clone)]
pub struct QNetwork {
weights: Vec<Vec<f64>>,
state_dim: usize,
action_dim: usize,
}
impl QNetwork {
pub fn new(state_dim: usize, action_dim: usize, hidden_dim: usize) -> Self {
use rand::RngExt;
let mut rng = rand::rng();
let mut weights = Vec::new();
let mut layer1 = Vec::new();
for _ in 0..(state_dim * hidden_dim) {
layer1.push(rng.random_range(-0.1..0.1));
}
weights.push(layer1);
let mut layer2 = Vec::new();
for _ in 0..(hidden_dim * action_dim) {
layer2.push(rng.random_range(-0.1..0.1));
}
weights.push(layer2);
Self {
weights,
state_dim,
action_dim,
}
}
pub fn forward(&self, state: &[f64]) -> Vec<f64> {
let mut output = vec![0.0; self.action_dim];
for (i, out) in output.iter_mut().enumerate().take(self.action_dim) {
for (j, &state_val) in state
.iter()
.enumerate()
.take(state.len().min(self.state_dim))
{
if let Some(layer) = self.weights.get(1) {
if let Some(&weight) = layer.get(j * self.action_dim + i) {
*out += state_val * weight;
}
}
}
}
output
}
pub fn update(&mut self, _learning_rate: f64) {
}
}
pub struct RLMarketMaker {
config: RLMarketMakerConfig,
q_network: QNetwork,
target_network: QNetwork,
replay_buffer: ReplayBuffer,
step_count: usize,
max_inventory: Decimal,
max_price: Decimal,
}
impl RLMarketMaker {
pub fn new(config: RLMarketMakerConfig, max_inventory: Decimal, max_price: Decimal) -> Self {
let state_dim = RLMarketState::dimension();
let action_dim = MarketAction::dimension();
let hidden_dim = 64;
let q_network = QNetwork::new(state_dim, action_dim, hidden_dim);
let target_network = q_network.clone();
let replay_buffer = ReplayBuffer::new(config.replay_buffer_capacity);
Self {
config,
q_network,
target_network,
replay_buffer,
step_count: 0,
max_inventory,
max_price,
}
}
pub fn select_action(&self, state: &RLMarketState, epsilon: f64) -> MarketAction {
use rand::RngExt;
let mut rng = rand::rng();
if rng.random_range(0.0..1.0) < epsilon {
MarketAction {
bid_offset: Decimal::from_f64_retain(rng.random_range(0.0..0.01)).unwrap(),
ask_offset: Decimal::from_f64_retain(rng.random_range(0.0..0.01)).unwrap(),
bid_quantity: Decimal::from(rng.random_range(50..200)),
ask_quantity: Decimal::from(rng.random_range(50..200)),
}
} else {
let state_vec = state.normalize(self.max_inventory, self.max_price);
let q_values = self.q_network.forward(&state_vec);
MarketAction::from_continuous(&q_values)
}
}
pub fn calculate_reward(
&self,
pnl: Decimal,
inventory: Decimal,
adverse_selection: Decimal,
) -> Decimal {
let inventory_cost = self.config.inventory_penalty * inventory.abs();
let adverse_cost = self.config.adverse_selection_penalty * adverse_selection;
pnl - inventory_cost - adverse_cost
}
pub fn store_experience(&mut self, experience: Experience) {
self.replay_buffer.push(experience);
}
pub fn train(&mut self) -> Result<Decimal> {
if self.replay_buffer.len() < self.config.batch_size {
return Ok(Decimal::ZERO);
}
let batch = self.replay_buffer.sample(self.config.batch_size);
let mut total_loss = 0.0;
for experience in &batch {
let state_vec = experience
.state
.normalize(self.max_inventory, self.max_price);
let next_state_vec = experience
.next_state
.normalize(self.max_inventory, self.max_price);
let q_values = self.q_network.forward(&state_vec);
let next_q_values = self.target_network.forward(&next_state_vec);
let max_next_q = next_q_values
.iter()
.fold(f64::NEG_INFINITY, |a, &b| a.max(b));
let reward_f64: f64 = experience.reward.to_string().parse().unwrap_or(0.0);
let target = if experience.done {
reward_f64
} else {
reward_f64 + self.config.discount_factor * max_next_q
};
for &q in &q_values {
total_loss += (q - target).powi(2);
}
}
self.q_network.update(self.config.learning_rate);
self.step_count += 1;
if self.step_count % self.config.target_update_freq == 0 {
self.target_network = self.q_network.clone();
}
Ok(Decimal::from_f64_retain(total_loss / batch.len() as f64).unwrap_or(Decimal::ZERO))
}
pub fn get_epsilon(&self) -> f64 {
self.config.epsilon.max(self.config.epsilon_min)
}
pub fn decay_epsilon(&mut self) {
self.config.epsilon *= self.config.epsilon_decay;
self.config.epsilon = self.config.epsilon.max(self.config.epsilon_min);
}
pub fn get_stats(&self) -> TrainingStats {
TrainingStats {
step_count: self.step_count,
epsilon: self.get_epsilon(),
replay_buffer_size: self.replay_buffer.len(),
}
}
pub fn save_model(&self, _path: &str) -> Result<()> {
Ok(())
}
pub fn load_model(&mut self, _path: &str) -> Result<()> {
Ok(())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TrainingStats {
pub step_count: usize,
pub epsilon: f64,
pub replay_buffer_size: usize,
}
pub struct OnlineLearningMarketMaker {
rl_maker: RLMarketMaker,
performance_window: VecDeque<Decimal>,
window_size: usize,
}
impl OnlineLearningMarketMaker {
pub fn new(config: RLMarketMakerConfig, max_inventory: Decimal, max_price: Decimal) -> Self {
let rl_maker = RLMarketMaker::new(config, max_inventory, max_price);
Self {
rl_maker,
performance_window: VecDeque::new(),
window_size: 100,
}
}
pub fn update_performance(&mut self, reward: Decimal) {
if self.performance_window.len() >= self.window_size {
self.performance_window.pop_front();
}
self.performance_window.push_back(reward);
}
pub fn get_avg_performance(&self) -> Decimal {
if self.performance_window.is_empty() {
return Decimal::ZERO;
}
let sum: Decimal = self.performance_window.iter().sum();
sum / Decimal::from(self.performance_window.len())
}
pub fn adapt_learning_rate(&mut self) {
let avg_perf = self.get_avg_performance();
if avg_perf < Decimal::ZERO {
self.rl_maker.config.learning_rate *= 1.1;
} else {
self.rl_maker.config.learning_rate *= 0.99;
}
self.rl_maker.config.learning_rate = self.rl_maker.config.learning_rate.clamp(0.0001, 0.01);
}
pub fn rl_maker(&self) -> &RLMarketMaker {
&self.rl_maker
}
pub fn rl_maker_mut(&mut self) -> &mut RLMarketMaker {
&mut self.rl_maker
}
}
#[cfg(test)]
mod tests {
use super::*;
use rust_decimal_macros::dec;
#[test]
fn test_market_state_normalize() {
let state = RLMarketState {
inventory: dec!(100),
mid_price: dec!(50000),
volatility: dec!(0.02),
spread: dec!(10),
order_book_imbalance: dec!(0.1),
recent_pnl: dec!(50),
time_elapsed: dec!(0.5),
trend: dec!(0.01),
};
let normalized = state.normalize(dec!(1000), dec!(100000));
assert_eq!(normalized.len(), RLMarketState::dimension());
for &val in &normalized {
assert!((-1.5..=1.5).contains(&val));
}
}
#[test]
fn test_market_action_from_continuous() {
let values = vec![0.005, 0.005, 0.5, 0.5];
let action = MarketAction::from_continuous(&values);
assert!(action.bid_offset >= Decimal::ZERO);
assert!(action.ask_offset >= Decimal::ZERO);
assert!(action.bid_quantity > Decimal::ZERO);
assert!(action.ask_quantity > Decimal::ZERO);
}
#[test]
fn test_replay_buffer() {
let mut buffer = ReplayBuffer::new(5);
let state = RLMarketState {
inventory: dec!(0),
mid_price: dec!(50000),
volatility: dec!(0.02),
spread: dec!(10),
order_book_imbalance: dec!(0),
recent_pnl: dec!(0),
time_elapsed: dec!(0),
trend: dec!(0),
};
let action = MarketAction {
bid_offset: dec!(0.001),
ask_offset: dec!(0.001),
bid_quantity: dec!(100),
ask_quantity: dec!(100),
};
for i in 0..10 {
buffer.push(Experience {
state: state.clone(),
action: action.clone(),
reward: Decimal::from(i),
next_state: state.clone(),
done: false,
});
}
assert_eq!(buffer.len(), 5);
let sample = buffer.sample(3);
assert_eq!(sample.len(), 3);
}
#[test]
fn test_rl_market_maker_creation() {
let config = RLMarketMakerConfig::default();
let maker = RLMarketMaker::new(config, dec!(1000), dec!(100000));
assert_eq!(maker.step_count, 0);
assert_eq!(maker.replay_buffer.len(), 0);
}
#[test]
fn test_select_action() {
let config = RLMarketMakerConfig::default();
let maker = RLMarketMaker::new(config, dec!(1000), dec!(100000));
let state = RLMarketState {
inventory: dec!(100),
mid_price: dec!(50000),
volatility: dec!(0.02),
spread: dec!(10),
order_book_imbalance: dec!(0.1),
recent_pnl: dec!(50),
time_elapsed: dec!(0.5),
trend: dec!(0.01),
};
let action = maker.select_action(&state, 1.0);
assert!(action.bid_offset >= Decimal::ZERO);
assert!(action.ask_offset >= Decimal::ZERO);
let action = maker.select_action(&state, 0.0);
assert!(action.bid_quantity > Decimal::ZERO);
assert!(action.ask_quantity > Decimal::ZERO);
}
#[test]
fn test_calculate_reward() {
let config = RLMarketMakerConfig::default();
let maker = RLMarketMaker::new(config, dec!(1000), dec!(100000));
let reward = maker.calculate_reward(dec!(100), dec!(50), dec!(10));
assert!(reward > dec!(99) && reward < dec!(100));
}
#[test]
fn test_epsilon_decay() {
let config = RLMarketMakerConfig::default();
let mut maker = RLMarketMaker::new(config, dec!(1000), dec!(100000));
let initial_epsilon = maker.get_epsilon();
maker.decay_epsilon();
let decayed_epsilon = maker.get_epsilon();
assert!(decayed_epsilon < initial_epsilon);
assert!(decayed_epsilon >= maker.config.epsilon_min);
}
#[test]
fn test_online_learning_market_maker() {
let config = RLMarketMakerConfig::default();
let mut maker = OnlineLearningMarketMaker::new(config, dec!(1000), dec!(100000));
maker.update_performance(dec!(10));
maker.update_performance(dec!(20));
maker.update_performance(dec!(30));
let avg = maker.get_avg_performance();
assert_eq!(avg, dec!(20));
}
#[test]
fn test_adapt_learning_rate() {
let config = RLMarketMakerConfig::default();
let mut maker = OnlineLearningMarketMaker::new(config, dec!(1000), dec!(100000));
let initial_lr = maker.rl_maker().config.learning_rate;
for _ in 0..10 {
maker.update_performance(dec!(-10));
}
maker.adapt_learning_rate();
let new_lr = maker.rl_maker().config.learning_rate;
assert!(new_lr > initial_lr);
}
#[test]
fn test_training_stats() {
let config = RLMarketMakerConfig::default();
let maker = RLMarketMaker::new(config, dec!(1000), dec!(100000));
let stats = maker.get_stats();
assert_eq!(stats.step_count, 0);
assert_eq!(stats.replay_buffer_size, 0);
assert!(stats.epsilon > 0.0);
}
}