feat(ml): BF16 VarBuilder, checkpoints, and training tensors for PPO

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-03 16:16:07 +01:00
parent d0dd3af2f1
commit 4e0225e090
15 changed files with 57 additions and 41 deletions

View File

@@ -6,11 +6,12 @@
use std::sync::Mutex;
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
use candle_nn::{VarBuilder, VarMap};
use crate::diffusion::config::DiffusionConfig;
use crate::diffusion::denoiser::Denoiser;
use crate::dqn::mixed_precision::training_dtype;
use crate::ensemble::inference_adapter::{
EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction,
};
@@ -46,7 +47,7 @@ impl DiffusionInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let data_dim = config.seq_len * config.feature_dim;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let denoiser = Denoiser::new(
data_dim,
config.hidden_dim,
@@ -69,7 +70,7 @@ impl DiffusionInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let data_dim = config.seq_len * config.feature_dim;
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let denoiser = Denoiser::new(
data_dim,
config.hidden_dim,

View File

@@ -5,9 +5,10 @@
use std::sync::Mutex;
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
use candle_nn::{VarBuilder, VarMap};
use crate::dqn::mixed_precision::training_dtype;
use crate::ensemble::inference_adapter::{
EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction,
};
@@ -41,7 +42,7 @@ impl KanInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let input_dim = config.layer_widths.first().copied().unwrap_or(51);
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let network = KANNetwork::new(&config, vb)?;
Ok(Self {
@@ -57,7 +58,7 @@ impl KanInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let input_dim = config.layer_widths.first().copied().unwrap_or(51);
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let network = KANNetwork::new(&config, vb)?;
varmap

View File

@@ -5,9 +5,10 @@
use std::sync::Mutex;
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
use candle_nn::{VarBuilder, VarMap};
use crate::dqn::mixed_precision::training_dtype;
use crate::ensemble::inference_adapter::{
EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta,
};
@@ -43,7 +44,7 @@ impl LiquidInferenceAdapter {
let device = config.device.resolve().unwrap_or(Device::Cpu);
let input_size = config.input_size;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let network = CandleCfCNetwork::new(&config, &vb)?;
Ok(Self {
@@ -59,7 +60,7 @@ impl LiquidInferenceAdapter {
let device = config.device.resolve().unwrap_or(Device::Cpu);
let input_size = config.input_size;
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let network = CandleCfCNetwork::new(&config, &vb)?;
varmap

View File

@@ -5,9 +5,10 @@
use std::sync::Mutex;
use candle_core::{DType, Device, Module, Tensor};
use candle_core::{Device, Module, Tensor};
use candle_nn::{linear, Linear, VarBuilder, VarMap};
use crate::dqn::mixed_precision::training_dtype;
use crate::ensemble::inference_adapter::{
EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction,
};
@@ -74,7 +75,7 @@ impl TggnInferenceAdapter {
pub fn new(input_dim: usize, hidden_dim: usize) -> MLResult<Self> {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let projection = TggnProjection::new(input_dim, hidden_dim, vb)?;
Ok(Self {
@@ -93,7 +94,7 @@ impl TggnInferenceAdapter {
) -> MLResult<Self> {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let projection = TggnProjection::new(input_dim, hidden_dim, vb)?;
varmap

View File

@@ -9,9 +9,10 @@
use std::collections::VecDeque;
use std::sync::Mutex;
use candle_core::{DType, Device, Module, Tensor};
use candle_core::{Device, Module, Tensor};
use candle_nn::{linear, Linear, VarBuilder, VarMap};
use crate::dqn::mixed_precision::training_dtype;
use crate::ensemble::inference_adapter::{
EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction,
};
@@ -97,7 +98,7 @@ impl TlobInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let flat_dim = sequence_length * feature_dim;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let projection = TlobProjection::new(flat_dim, hidden_dim, vb)?;
Ok(Self {
@@ -120,7 +121,7 @@ impl TlobInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let flat_dim = sequence_length * feature_dim;
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let projection = TlobProjection::new(flat_dim, hidden_dim, vb)?;
varmap

View File

@@ -8,9 +8,10 @@
use std::collections::VecDeque;
use std::sync::Mutex;
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
use candle_nn::{VarBuilder, VarMap};
use crate::dqn::mixed_precision::training_dtype;
use crate::ensemble::inference_adapter::{
EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction,
};
@@ -54,7 +55,7 @@ impl XlstmInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let input_dim = config.input_dim;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let network = XLSTMNetwork::new(&config, vb)?;
Ok(Self {
@@ -76,7 +77,7 @@ impl XlstmInferenceAdapter {
let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu);
let input_dim = config.input_dim;
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
let network = XLSTMNetwork::new(&config, vb)?;
varmap

View File

@@ -168,7 +168,7 @@ impl IntegratedGradients {
#[cfg(test)]
mod tests {
use super::*;
use candle_core::{DType, Device};
use candle_core::Device;
use candle_nn::{linear, VarBuilder, VarMap};
/// A simple 2-layer test network: linear(4,8) -> relu -> linear(8,1)
@@ -197,7 +197,7 @@ mod tests {
fn test_integrated_gradients_basic() {
let device = Device::Cpu;
let varmap = VarMap::new();
let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vs = VarBuilder::from_varmap(&varmap, crate::dqn::mixed_precision::training_dtype(&device), &device);
let model = TwoLayerNet::new(vs);
// Linear layers require 2D input: [batch, features]
@@ -244,7 +244,7 @@ mod tests {
fn test_ig_completeness_axiom() {
let device = Device::Cpu;
let varmap = VarMap::new();
let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vs = VarBuilder::from_varmap(&varmap, crate::dqn::mixed_precision::training_dtype(&device), &device);
let model = LinearNet::new(vs);
let input = Tensor::new(&[[1.0_f32, -0.5, 0.3, 2.0]], &device).unwrap();
@@ -298,7 +298,7 @@ mod tests {
fn test_ig_dimension_mismatch() {
let device = Device::Cpu;
let varmap = VarMap::new();
let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vs = VarBuilder::from_varmap(&varmap, crate::dqn::mixed_precision::training_dtype(&device), &device);
let model = TwoLayerNet::new(vs);
let input = Tensor::new(&[[1.0_f32, -0.5, 0.3, 2.0]], &device).unwrap();

View File

@@ -18,6 +18,7 @@ use candle_core::{DType, Device, Tensor};
use candle_nn::{linear, Linear, Module, VarBuilder, VarMap};
use super::bar_resampler::BarResampler;
use crate::dqn::mixed_precision::training_dtype;
use crate::types::OHLCVBar;
use crate::MLError;
@@ -266,7 +267,7 @@ impl MultiTimeframeEncoder {
device: &Device,
) -> Result<(Self, VarMap), MLError> {
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, DType::F32, device);
let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device);
let encoder = Self::new(config, vb)?;
Ok((encoder, vars))
}
@@ -557,7 +558,7 @@ mod tests {
#[test]
fn test_lstm_encoder_single_step() {
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, DType::F32, &Device::Cpu);
let vb = VarBuilder::from_varmap(&vars, training_dtype(&Device::Cpu), &Device::Cpu);
let lstm = LstmEncoder::new(6, 32, vb.pp("test_lstm")).expect("lstm creation");
// Single timestep: (1, 6)
@@ -573,7 +574,7 @@ mod tests {
#[test]
fn test_lstm_encoder_multi_step() {
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, DType::F32, &Device::Cpu);
let vb = VarBuilder::from_varmap(&vars, training_dtype(&Device::Cpu), &Device::Cpu);
let lstm = LstmEncoder::new(6, 64, vb.pp("test_lstm")).expect("lstm creation");
// 10 timesteps: (10, 6)

View File

@@ -4,13 +4,14 @@
//! in high-frequency trading. Unlike traditional time-series transformers, this model
//! operates directly on portfolio state vectors for optimal weight prediction.
use candle_core::{DType, Device, IndexOp, Module, ModuleT, Result as CandleResult, Tensor};
use candle_core::{Device, IndexOp, Module, ModuleT, Result as CandleResult, Tensor};
use candle_nn::{Linear, VarBuilder, VarMap};
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use tracing::{debug, instrument, warn};
use super::*;
use crate::dqn::mixed_precision::training_dtype;
/// Portfolio state representation for transformer input
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -187,7 +188,7 @@ impl PortfolioTransformer {
/// Create new Portfolio Transformer
pub fn new(config: PortfolioTransformerConfig, device: Device) -> MLResult<Self> {
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device);
// Input projection
let input_projection = candle_nn::linear(

View File

@@ -12,7 +12,7 @@
use std::f32::consts::PI;
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
use candle_nn::{linear, Linear, Module, VarBuilder, VarMap};
use rand::thread_rng;
// Note: rand_distr::Distribution could be used for direct sampling if added to dependencies
@@ -21,6 +21,7 @@ use serde::{Deserialize, Serialize};
use statrs::distribution::{ContinuousCDF, Normal};
use tracing::{debug, warn};
use crate::dqn::mixed_precision::training_dtype;
use crate::dqn::xavier_init::linear_xavier;
use crate::MLError;
@@ -80,7 +81,7 @@ impl ContinuousPolicyNetwork {
/// Create new continuous policy network
pub fn new(config: ContinuousPolicyConfig, device: Device) -> Result<Self, MLError> {
let vars = VarMap::new();
let var_builder = VarBuilder::from_varmap(&vars, DType::F32, &device);
let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device);
let mut feature_layers = Vec::new();
let mut current_dim = config.state_dim;

View File

@@ -9,6 +9,7 @@ use rand::thread_rng;
use rand_distr::{Distribution, Normal};
use serde::{Deserialize, Serialize};
use crate::dqn::mixed_precision::training_dtype;
use crate::dqn::xavier_init::linear_xavier;
use crate::MLError;
@@ -121,7 +122,7 @@ impl FlowPolicy {
/// A new FlowPolicy instance or an error if initialization fails.
pub fn new(config: FlowPolicyConfig, device: &Device) -> Result<Self, MLError> {
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, DType::F32, device);
let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device);
// Context encoder: state_dim → context_dim
let context_enc = linear_xavier(

View File

@@ -5,10 +5,11 @@
//! these networks maintain hidden states across timesteps, enabling the agent to
//! "remember" past observations when making decisions.
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
use candle_nn::{linear, Linear, Module, VarBuilder, VarMap, LSTM, LSTMConfig};
use candle_nn::rnn::{RNN, LSTMState};
use crate::dqn::mixed_precision::training_dtype;
use crate::MLError;
/// LSTM-augmented policy network for temporal action selection
@@ -49,7 +50,7 @@ impl LSTMPolicyNetwork {
device: Device,
) -> Result<Self, MLError> {
let vars = VarMap::new();
let var_builder = VarBuilder::from_varmap(&vars, DType::F32, &device);
let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device);
// Input projection layer (state_dim → hidden_dim)
let input_layer = linear(input_dim, hidden_dim, var_builder.pp("input"))
@@ -264,7 +265,7 @@ impl LSTMValueNetwork {
device: Device,
) -> Result<Self, MLError> {
let vars = VarMap::new();
let var_builder = VarBuilder::from_varmap(&vars, DType::F32, &device);
let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device);
// Input projection layer (state_dim → hidden_dim)
let input_layer = linear(input_dim, hidden_dim, var_builder.pp("input"))

View File

@@ -28,6 +28,7 @@ use super::hidden_state_manager::HiddenStateManager;
use super::lstm_networks::{LSTMPolicyNetwork, LSTMValueNetwork};
use super::trajectories::{TrajectoryBatch, TrajectoryTensors};
use crate::dqn::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig};
use crate::dqn::mixed_precision::training_dtype;
use crate::dqn::portfolio_tracker::PortfolioTracker;
use crate::dqn::xavier_init::linear_xavier;
use crate::dqn::reward::RewardNormalizer;
@@ -298,7 +299,7 @@ impl PolicyNetwork {
device: Device,
) -> Result<Self, MLError> {
let vars = VarMap::new();
let var_builder = VarBuilder::from_varmap(&vars, DType::F32, &device);
let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device);
let mut layers = Vec::new();
let mut current_dim = input_dim;
@@ -546,7 +547,7 @@ impl ValueNetwork {
/// Create new value network
pub fn new(input_dim: usize, hidden_dims: &[usize], device: Device) -> Result<Self, MLError> {
let vars = VarMap::new();
let var_builder = VarBuilder::from_varmap(&vars, DType::F32, &device);
let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device);
let mut layers = Vec::new();
let mut current_dim = input_dim;
@@ -1792,7 +1793,7 @@ impl PPO {
// 4. Candle's deserializer validates format before tensor creation
// 5. Any format violations cause Err return, not UB
let actor_vb = unsafe {
VarBuilder::from_mmaped_safetensors(&[actor_path], DType::F32, &device).map_err(
VarBuilder::from_mmaped_safetensors(&[actor_path], training_dtype(&device), &device).map_err(
|e| {
MLError::ModelError(format!(
"Failed to load actor checkpoint from {}: {}",
@@ -1849,7 +1850,7 @@ impl PPO {
// 4. Candle's deserializer validates format before tensor creation
// 5. Any format violations cause Err return, not UB
let critic_vb = unsafe {
VarBuilder::from_mmaped_safetensors(&[critic_path], DType::F32, &device).map_err(
VarBuilder::from_mmaped_safetensors(&[critic_path], training_dtype(&device), &device).map_err(
|e| {
MLError::ModelError(format!(
"Failed to load critic checkpoint from {}: {}",

View File

@@ -555,7 +555,6 @@ impl OnlineLearner {
#[cfg(test)]
mod tests {
use super::*;
use candle_core::DType;
// -----------------------------------------------------------------------
// Helpers
@@ -578,7 +577,7 @@ mod tests {
/// Build a tiny VarMap with a single 2x2 parameter for testing.
fn tiny_var_map() -> Result<VarMap, MLError> {
let var_map = VarMap::new();
let vb = candle_nn::VarBuilder::from_varmap(&var_map, DType::F32, &Device::Cpu);
let vb = candle_nn::VarBuilder::from_varmap(&var_map, crate::dqn::mixed_precision::training_dtype(&Device::Cpu), &Device::Cpu);
let _linear = candle_nn::linear(2, 2, vb.pp("layer"))
.map_err(|e| MLError::ModelError(format!("tiny_var_map linear: {e}")))?;
Ok(var_map)

View File

@@ -25,6 +25,7 @@ use crate::ppo::gae::GAEConfig;
use crate::ppo::ppo::{PPOConfig, PPO};
use crate::ppo::trajectories::{Trajectory, TrajectoryBatch, TrajectoryStep};
use crate::cuda_pipeline::PpoGpuData;
use crate::dqn::mixed_precision::training_dtype;
use crate::MLError;
/// PPO training hyperparameters (matches gRPC PpoParams)
@@ -832,6 +833,7 @@ impl PpoTrainer {
.map_err(|e| MLError::ModelError(format!("PPO GPU state access failed: {e}")))?
} else {
Tensor::from_vec(state.clone(), (1, state.len()), &self.device)?
.to_dtype(training_dtype(&self.device))?
};
// Actor forward pass — use pre-uploaded tensor
@@ -892,7 +894,8 @@ impl PpoTrainer {
// Phase 2: Batched critic forward — SINGLE GPU→CPU sync for all value estimates
if step_count > 0 && state_dim > 0 {
let all_states_tensor =
Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)?;
Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)?
.to_dtype(training_dtype(&self.device))?;
let all_values_vec = model
.critic
.forward(&all_states_tensor)?
@@ -958,6 +961,7 @@ impl PpoTrainer {
.map_err(|e| MLError::ModelError(format!("PPO GPU state access failed: {e}")))?
} else {
Tensor::from_vec(state.clone(), (1, state.len()), &self.device)?
.to_dtype(training_dtype(&self.device))?
};
// Get action from policy — use pre-uploaded tensor
@@ -1011,7 +1015,8 @@ impl PpoTrainer {
// Phase 2: Batched critic forward — SINGLE GPU→CPU sync for all value estimates
if step_count > 0 && state_dim > 0 {
let all_states_tensor =
Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)?;
Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)?
.to_dtype(training_dtype(&self.device))?;
let all_values_vec = model
.critic
.forward(&all_states_tensor)?