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:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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 {}: {}",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)?
|
||||
|
||||
Reference in New Issue
Block a user