feat(ml): factored 45-action masking, confidence EMA, RAII memory docs

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-23 01:03:44 +01:00
parent 4e228f127f
commit cfadd7d3dc
4 changed files with 215 additions and 23 deletions

View File

@@ -145,6 +145,10 @@ pub struct SecondaryBettingModel {
total_predictions: Arc<AtomicU64>,
total_trades: Arc<AtomicU64>,
total_bet_size: Arc<AtomicU64>, // Stored as fixed-point (multiply by 1e6)
/// Exponential moving average of prediction confidence, fixed-point (x 1_000_000)
confidence_ema: Arc<AtomicU64>,
/// EMA smoothing factor (higher = more weight on recent observations)
confidence_ema_alpha: f64,
}
impl SecondaryBettingModel {
@@ -157,6 +161,8 @@ impl SecondaryBettingModel {
total_predictions: Arc::new(AtomicU64::new(0)),
total_trades: Arc::new(AtomicU64::new(0)),
total_bet_size: Arc::new(AtomicU64::new(0)),
confidence_ema: Arc::new(AtomicU64::new(0)),
confidence_ema_alpha: 0.99,
})
}
@@ -188,6 +194,21 @@ impl SecondaryBettingModel {
// Update statistics
self.total_predictions.fetch_add(1, Ordering::Relaxed);
// Update confidence EMA (fixed-point encoding: value * 1_000_000)
{
let current_ema_fixed = self.confidence_ema.load(Ordering::Relaxed);
let current_ema = current_ema_fixed as f64 / 1_000_000.0;
let alpha = self.confidence_ema_alpha;
let new_ema = if current_ema_fixed == 0 && self.total_predictions.load(Ordering::Relaxed) == 1 {
// First prediction: initialize EMA to this confidence value
primary.confidence
} else {
alpha * current_ema + (1.0 - alpha) * primary.confidence
};
self.confidence_ema
.store((new_ema * 1_000_000.0) as u64, Ordering::Relaxed);
}
// Validate inputs
if primary.features.is_empty() || features.is_empty() {
return Err(MLError::ValidationError {
@@ -335,7 +356,7 @@ impl SecondaryBettingModel {
total_trades,
total_rejections: total_predictions.saturating_sub(total_trades),
average_bet_size,
average_confidence: 0.0, // TODO: Track confidence running average
average_confidence: self.confidence_ema.load(Ordering::Relaxed) as f64 / 1_000_000.0,
}
}
@@ -344,6 +365,7 @@ impl SecondaryBettingModel {
self.total_predictions.store(0, Ordering::Relaxed);
self.total_trades.store(0, Ordering::Relaxed);
self.total_bet_size.store(0, Ordering::Relaxed);
self.confidence_ema.store(0, Ordering::Relaxed);
}
}
@@ -394,6 +416,42 @@ mod tests {
assert!(combined < 0.6);
}
#[test]
fn test_secondary_model_confidence_ema() {
let config = SecondaryModelConfig::default();
let mut model = SecondaryBettingModel::new(config).unwrap();
let primary = PrimaryPrediction {
direction: 1,
confidence: 0.8,
expected_return: 0.05,
features: vec![0.5, 0.6, 0.7],
};
let features = vec![0.2, 0.8, 0.7];
// First prediction initializes EMA to 0.8
let _ = model.should_trade(&primary, &features).unwrap();
let stats = model.get_statistics();
assert!((stats.average_confidence - 0.8).abs() < 0.01);
// Second prediction with confidence 0.6
// EMA = 0.99 * 0.8 + 0.01 * 0.6 = 0.798
let primary2 = PrimaryPrediction {
direction: 1,
confidence: 0.6,
expected_return: 0.05,
features: vec![0.5, 0.6, 0.7],
};
let _ = model.should_trade(&primary2, &features).unwrap();
let stats2 = model.get_statistics();
assert!((stats2.average_confidence - 0.798).abs() < 0.01);
// Reset clears EMA
model.reset_statistics();
let stats3 = model.get_statistics();
assert!((stats3.average_confidence - 0.0).abs() < 1e-6);
}
#[test]
fn test_bet_size_calculation() {
let config = SecondaryModelConfig::default();

View File

@@ -76,14 +76,32 @@ pub fn create_action_mask(current_position: f64, max_position: f64, num_actions:
}
// HOLD (action 2) is always valid
} else {
// Future 45-action implementation (Phase 3)
// This will be implemented when expanding to factored action space
// For now, default all to valid
// TODO(Phase 3): Implement factored action masking
// - Map action index → FactoredAction
// - Get target_exposure from action
// - Mask if |target_exposure| > max_position
} else if num_actions == 45 {
// Factored 45-action space: (direction: 3) x (size: 5) x (urgency: 3) = 45
// direction = action_idx / 15: 0=Buy, 1=Sell, 2=Hold
for action_idx in 0..num_actions {
let direction = action_idx / 15;
match direction {
0 => {
// Buy: mask if at or above max position
if current_position >= max_position {
if let Some(m) = mask.get_mut(action_idx) {
*m = false;
}
}
}
1 => {
// Sell: mask if at or below negative max position
if current_position <= -max_position {
if let Some(m) = mask.get_mut(action_idx) {
*m = false;
}
}
}
// 2 = Hold: always valid, no masking needed
_ => {}
}
}
}
mask
@@ -192,6 +210,70 @@ mod tests {
assert_eq!(values[2], 0.2);
}
#[test]
fn test_create_action_mask_45_flat_position() {
// Flat position: all 45 actions valid
let mask = create_action_mask(0.0, 2.0, 45);
assert_eq!(mask.len(), 45);
assert!(mask.iter().all(|&v| v));
}
#[test]
fn test_create_action_mask_45_at_max_position() {
// At max long: Buy actions (0..15) masked, Sell (15..30) and Hold (30..45) valid
let mask = create_action_mask(2.0, 2.0, 45);
assert_eq!(mask.len(), 45);
// Buy actions (direction=0, indices 0..15) should be masked
for i in 0..15 {
assert_eq!(
mask.get(i).copied().unwrap_or(true),
false,
"Buy action {} should be masked at max position",
i
);
}
// Sell actions (direction=1, indices 15..30) should be valid
for i in 15..30 {
assert_eq!(
mask.get(i).copied().unwrap_or(false),
true,
"Sell action {} should be valid at max position",
i
);
}
// Hold actions (direction=2, indices 30..45) should be valid
for i in 30..45 {
assert_eq!(
mask.get(i).copied().unwrap_or(false),
true,
"Hold action {} should be valid at max position",
i
);
}
}
#[test]
fn test_create_action_mask_45_at_min_position() {
// At max short: Sell actions (15..30) masked, Buy (0..15) and Hold (30..45) valid
let mask = create_action_mask(-2.0, 2.0, 45);
for i in 0..15 {
assert_eq!(mask.get(i).copied().unwrap_or(false), true);
}
for i in 15..30 {
assert_eq!(mask.get(i).copied().unwrap_or(true), false);
}
for i in 30..45 {
assert_eq!(mask.get(i).copied().unwrap_or(false), true);
}
}
#[test]
fn test_create_action_mask_45_partial_position() {
// Partial position: neither at max nor min, all valid
let mask = create_action_mask(1.0, 2.0, 45);
assert!(mask.iter().all(|&v| v));
}
#[test]
fn test_apply_mask_all_invalid() {
let device = Device::Cpu;

View File

@@ -22,7 +22,7 @@
//! - Penalty = penalty_coeff × max(0, |action| - soft_threshold)²
//! - Encourages staying away from hard limits (margin of safety)
use candle_core::{Device, Tensor};
use candle_core::{Device, IndexOp, Tensor};
use serde::{Deserialize, Serialize};
use crate::MLError;
@@ -90,12 +90,22 @@ impl ContinuousActionConstraints {
)));
}
// TODO(Phase 2): Extract current position from state to compute dynamic bounds
// Example: let current_position = state.get(0)?.get(0)?.to_scalar::<f32>()?;
// Extract current position from state tensor (first element of first sample).
// Safe access: returns 0.0 if tensor indexing fails.
let current_position = state
.i((0, 0))
.and_then(|t| t.to_scalar::<f32>())
.unwrap_or(0.0);
// Compute asymmetric bounds based on current position:
// - max_long: remaining room to go long
// - max_short: remaining room to go short
let max_long = max_position_abs - current_position.max(0.0);
let max_short = max_position_abs + current_position.min(0.0);
Ok(Self {
min_position: -max_position_abs,
max_position: max_position_abs,
min_position: -max_short,
max_position: max_long,
penalty_coeff: 1.0, // Moderate penalty
soft_threshold_fraction: 0.8, // Penalize beyond 80% of limit
})
@@ -570,6 +580,54 @@ mod tests {
assert!(!constraints.is_within_soft_bounds(2.0));
}
#[test]
fn test_from_state_asymmetric_bounds_long() {
let device = Device::Cpu;
// State with current_position = 1.5 in first element
let mut state_data = vec![0.0f32; 64];
state_data[0] = 1.5;
let state = Tensor::new(state_data.as_slice(), &device)
.unwrap()
.reshape((1, 64))
.unwrap();
let constraints = ContinuousActionConstraints::from_state(&state, 2.0).unwrap();
// max_long = 2.0 - max(1.5, 0.0) = 0.5
assert!((constraints.max_position - 0.5).abs() < 1e-6);
// max_short = 2.0 + min(1.5, 0.0) = 2.0, so min_position = -2.0
assert!((constraints.min_position - (-2.0)).abs() < 1e-6);
}
#[test]
fn test_from_state_asymmetric_bounds_short() {
let device = Device::Cpu;
// State with current_position = -1.0
let mut state_data = vec![0.0f32; 64];
state_data[0] = -1.0;
let state = Tensor::new(state_data.as_slice(), &device)
.unwrap()
.reshape((1, 64))
.unwrap();
let constraints = ContinuousActionConstraints::from_state(&state, 2.0).unwrap();
// max_long = 2.0 - max(-1.0, 0.0) = 2.0
assert!((constraints.max_position - 2.0).abs() < 1e-6);
// max_short = 2.0 + min(-1.0, 0.0) = 1.0, so min_position = -1.0
assert!((constraints.min_position - (-1.0)).abs() < 1e-6);
}
#[test]
fn test_from_state_flat_position_symmetric() {
let device = Device::Cpu;
// State with current_position = 0.0 (flat)
let state = Tensor::zeros((1, 64), DType::F32, &device).unwrap();
let constraints = ContinuousActionConstraints::from_state(&state, 2.0).unwrap();
// Symmetric: max_long = 2.0, max_short = 2.0
assert!((constraints.max_position - 2.0).abs() < 1e-6);
assert!((constraints.min_position - (-2.0)).abs() < 1e-6);
}
#[test]
fn test_mask_continuous_actions_integration() {
let device = Device::Cpu;

View File

@@ -358,15 +358,9 @@ impl SafeMemoryManager {
usage.reset_peak();
}
// Force garbage collection hint (if applicable)
// Note: Rust does not have a standard garbage collector
// This is a placeholder for future integration with alternative GC implementations
#[cfg(feature = "gc")]
{
// TODO: Integrate with a Rust GC library like `gc` or `rust-gc` if needed
// For now, this is a no-op as Rust uses RAII and ownership for memory management
tracing::debug!("GC hint requested but no GC is available in standard Rust");
}
// No garbage collector needed: Rust's ownership model and RAII handle
// deallocation automatically when Tensors and VarMaps go out of scope.
// GPU memory (CUDA/Metal) is freed when candle Tensors are dropped.
info!("Memory cleanup completed for device: {}", device_key);
Ok(())