From d95e205d4b3d43967bc1bba62686b60c7b732261 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 16 Mar 2026 16:11:48 +0100 Subject: [PATCH] =?UTF-8?q?refactor(ml):=20delete=20mixed=5Fprecision=20mo?= =?UTF-8?q?dule=20=E2=80=94=20BF16=20unconditional=20on=20CUDA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Eliminate the entire mixed_precision runtime indirection layer: - Delete crates/ml-core/src/mixed_precision.rs (training_dtype, ensure_training_dtype, align_dim_for_tensor_cores) - Inline ~100 call sites across 130 files to constants: training_dtype(&device) → candle_core::DType::BF16 ensure_training_dtype(x) → x.to_dtype(candle_core::DType::BF16) align_dim_for_tensor_cores(x, &device) → (x + 7) & !7 - Remove re-exports from ml-dqn, ml-supervised, ml lib.rs - Clean config/toml/json/shell references No CPU/Metal training path exists — BF16 is the only dtype. Co-Authored-By: Claude Opus 4.6 --- config/ml/inference.toml | 1 - config/ml/model_params.toml | 2 +- config/ml/training.toml | 1 - crates/ml-core/src/common/config.rs | 2 - crates/ml-core/src/lib.rs | 1 - crates/ml-core/src/mixed_precision.rs | 745 ------------------ crates/ml-dqn/src/agent.rs | 55 +- crates/ml-dqn/src/attention.rs | 15 +- crates/ml-dqn/src/branching.rs | 14 +- crates/ml-dqn/src/curiosity.rs | 13 +- crates/ml-dqn/src/distributional_dueling.rs | 5 +- crates/ml-dqn/src/dqn.rs | 86 +- crates/ml-dqn/src/dueling.rs | 5 +- crates/ml-dqn/src/gpu_replay_buffer.rs | 3 +- crates/ml-dqn/src/lib.rs | 1 - crates/ml-dqn/src/network.rs | 41 +- crates/ml-dqn/src/noisy_layers.rs | 32 +- crates/ml-dqn/src/quantile_regression.rs | 7 +- crates/ml-dqn/src/rainbow_agent.rs | 5 +- crates/ml-dqn/src/rainbow_network.rs | 5 +- crates/ml-dqn/src/replay_buffer_type.rs | 73 +- crates/ml-dqn/src/residual.rs | 33 +- crates/ml-dqn/src/rmsnorm.rs | 33 +- crates/ml-dqn/tests/gpu_smoketest.rs | 3 +- .../src/integrated_gradients.rs | 6 +- crates/ml-ppo/src/continuous_demo.rs | 6 +- crates/ml-ppo/src/continuous_policy.rs | 14 +- crates/ml-ppo/src/continuous_ppo.rs | 11 +- crates/ml-ppo/src/flow_policy/mod.rs | 15 +- crates/ml-ppo/src/hidden_state_manager.rs | 18 +- crates/ml-ppo/src/lstm_networks.rs | 9 +- crates/ml-ppo/src/ppo.rs | 87 +- crates/ml-ppo/src/trajectories.rs | 5 +- .../ml-supervised/src/diffusion/denoiser.rs | 16 +- crates/ml-supervised/src/diffusion/sampler.rs | 6 +- crates/ml-supervised/src/kan/layer.rs | 5 +- crates/ml-supervised/src/kan/network.rs | 7 +- crates/ml-supervised/src/lib.rs | 1 - crates/ml-supervised/src/liquid/candle_cfc.rs | 29 +- crates/ml-supervised/src/liquid/training.rs | 3 +- crates/ml-supervised/src/mamba/mod.rs | 5 +- crates/ml-supervised/src/mamba/ssd_layer.rs | 2 +- crates/ml-supervised/src/tft/mod.rs | 15 +- .../src/tft/quantized_attention.rs | 3 +- crates/ml-supervised/src/tft/quantized_grn.rs | 5 +- .../ml-supervised/src/tft/quantized_lstm.rs | 6 +- crates/ml-supervised/src/tft/quantized_vsn.rs | 6 +- .../src/tft/varmap_quantization.rs | 5 +- crates/ml-supervised/src/xlstm/block.rs | 9 +- crates/ml-supervised/src/xlstm/mlstm.rs | 13 +- crates/ml-supervised/src/xlstm/network.rs | 17 +- crates/ml-supervised/src/xlstm/slstm.rs | 11 +- crates/ml/examples/train_baseline_rl.rs | 4 +- crates/ml/src/benchmark/dqn_benchmark.rs | 18 +- crates/ml/src/benchmark/tft_benchmark.rs | 23 - crates/ml/src/cuda_pipeline/gpu_weights.rs | 4 +- crates/ml/src/cuda_pipeline/mod.rs | 11 +- crates/ml/src/diffusion/trainable.rs | 6 +- crates/ml/src/ensemble/adapters/diffusion.rs | 12 +- crates/ml/src/ensemble/adapters/dqn.rs | 5 +- crates/ml/src/ensemble/adapters/kan.rs | 12 +- crates/ml/src/ensemble/adapters/liquid.rs | 5 +- crates/ml/src/ensemble/adapters/mamba2.rs | 3 +- crates/ml/src/ensemble/adapters/ppo.rs | 5 +- crates/ml/src/ensemble/adapters/tft.rs | 6 +- crates/ml/src/ensemble/adapters/tggn.rs | 7 +- crates/ml/src/ensemble/adapters/tlob.rs | 7 +- crates/ml/src/ensemble/adapters/xlstm.rs | 10 +- crates/ml/src/features/mod.rs | 2 +- crates/ml/src/features/multi_timeframe.rs | 7 +- crates/ml/src/flash_attention/mod.rs | 31 - crates/ml/src/hyperopt/adapters/dqn.rs | 11 - crates/ml/src/hyperopt/adapters/ppo.rs | 5 - crates/ml/src/hyperopt/adapters/tft.rs | 2 - crates/ml/src/kan/trainable.rs | 6 +- crates/ml/src/lib.rs | 1 - crates/ml/src/liquid/adapter.rs | 4 +- crates/ml/src/portfolio_transformer.rs | 4 +- crates/ml/src/ppo/trainable_adapter.rs | 2 +- crates/ml/src/tft/training.rs | 3 - crates/ml/src/tgnn/trainable_adapter.rs | 6 +- crates/ml/src/tlob/trainable_adapter.rs | 6 +- crates/ml/src/trainers/dqn/config.rs | 126 +-- crates/ml/src/trainers/dqn/fused_training.rs | 115 --- .../dqn/smoke_tests/feature_coverage.rs | 4 +- .../trainers/dqn/smoke_tests/gpu_residency.rs | 16 +- .../src/trainers/dqn/smoke_tests/helpers.rs | 6 +- crates/ml/src/trainers/dqn/trainer/action.rs | 7 +- .../src/trainers/dqn/trainer/constructor.rs | 101 +-- crates/ml/src/trainers/dqn/trainer/metrics.rs | 65 +- crates/ml/src/trainers/dqn/trainer/state.rs | 2 +- crates/ml/src/trainers/dqn/trainer/tests.rs | 7 +- .../ml/src/trainers/dqn/trainer/train_step.rs | 24 +- .../src/trainers/dqn/trainer/training_loop.rs | 4 +- crates/ml/src/trainers/online_learning.rs | 2 +- crates/ml/src/trainers/ppo.rs | 13 - crates/ml/src/trainers/tft/config.rs | 1 - crates/ml/src/trainers/tlob.rs | 3 +- crates/ml/src/training/orchestrator.rs | 3 - crates/ml/src/training_pipeline.rs | 3 - crates/ml/src/validation/ppo_adapter.rs | 10 +- crates/ml/src/xlstm/trainable.rs | 6 +- .../dqn_accumulation_convergence_test.rs | 2 + .../tests/dqn_gradient_accumulation_test.rs | 1 + crates/ml/tests/dqn_inference_test.rs | 1 + crates/ml/tests/dqn_long_training_test.rs | 1 + crates/ml/tests/dqn_training_pipeline_test.rs | 6 + crates/ml/tests/dqn_training_smoke_test.rs | 1 + crates/ml/tests/gpu_per_integration_test.rs | 10 +- .../ml/tests/ppo_45_action_network_tests.rs | 1 - .../tests/ppo_recurrent_integration_tests.rs | 1 - .../tests/tft_inference_latency_benchmark.rs | 1 - crates/ml/tests/tft_real_dbn_data_test.rs | 1 - crates/ml/tests/tft_test.rs | 3 - .../plans/2026-03-15-fused-cuda-training.md | 259 ++++++ .../2026-03-15-fused-cuda-training-design.md | 416 ++++++++++ scripts/gpu-hotpath-guard.sh | 1 - .../src/ensemble_training_coordinator.rs | 1 - .../ml_training_service/src/gpu_config.rs | 17 - .../tests/ensemble_training_tests.rs | 1 - .../job_configs/invalid_zero_epochs.json | 1 - .../tests/fixtures/job_configs/valid_dqn.json | 1 - .../end_to_end_batch_workflow_test.rs | 1 - .../integration/failure_recovery_test.rs | 1 - .../integration/grpc_api_integration_test.rs | 1 - .../integration/real_data_integration_test.rs | 2 - .../tests/integration_tests.rs | 1 - .../tests/orchestrator_comprehensive_tests.rs | 1 - .../tests/training_error_recovery_tests.rs | 1 - .../tests/validation_pipeline_tests.rs | 1 - .../trading_service/src/services/ppo_model.rs | 1 - .../trading_service/src/services/tft_model.rs | 2 - 132 files changed, 1139 insertions(+), 1861 deletions(-) delete mode 100644 crates/ml-core/src/mixed_precision.rs create mode 100644 docs/superpowers/plans/2026-03-15-fused-cuda-training.md create mode 100644 docs/superpowers/specs/2026-03-15-fused-cuda-training-design.md diff --git a/config/ml/inference.toml b/config/ml/inference.toml index eb3930c37..5bbdda4b6 100644 --- a/config/ml/inference.toml +++ b/config/ml/inference.toml @@ -14,7 +14,6 @@ max_concurrent_requests = 1000 # Maximum concurrent inference requests # GPU configuration device_id = 0 # CUDA device ID memory_pool_mb = 1024 # GPU memory pool size in MB -enable_mixed_precision = true # Enable mixed precision inference cuda_streams = 4 # Number of CUDA streams tensor_rt_optimization = false # Enable TensorRT optimization diff --git a/config/ml/model_params.toml b/config/ml/model_params.toml index 20f032c18..28b1c21cc 100644 --- a/config/ml/model_params.toml +++ b/config/ml/model_params.toml @@ -90,7 +90,7 @@ warmup_iterations = 100 # Model warmup iterations on startup enable_gpu = true # Enable GPU acceleration device_id = 0 # CUDA device ID memory_pool_mb = 1024 # GPU memory pool size -enable_mixed_precision = true # Enable mixed precision inference + [model_paths] # Model file paths diff --git a/config/ml/training.toml b/config/ml/training.toml index 6041230ee..63a746269 100644 --- a/config/ml/training.toml +++ b/config/ml/training.toml @@ -135,7 +135,6 @@ gradient_compression = false # Enable gradient compression [memory_optimization] # Memory optimization gradient_checkpointing = false # Enable gradient checkpointing -mixed_precision = true # Enable mixed precision training memory_efficient_attention = true # Enable memory efficient attention cpu_offload = false # Enable CPU offloading diff --git a/crates/ml-core/src/common/config.rs b/crates/ml-core/src/common/config.rs index b30a37fe6..2c69ec693 100644 --- a/crates/ml-core/src/common/config.rs +++ b/crates/ml-core/src/common/config.rs @@ -127,7 +127,6 @@ pub struct HardwareConfig { pub use_gpu: bool, pub gpu_memory_limit_mb: Option, pub cpu_threads: Option, - pub enable_mixed_precision: bool, } impl MLConfig { @@ -228,7 +227,6 @@ impl HardwareConfig { use_gpu: false, // CPU only for safety gpu_memory_limit_mb: None, cpu_threads: Some(1), // Single thread to prevent resource issues - enable_mixed_precision: true, // Auto-detected per GPU capabilities } } } diff --git a/crates/ml-core/src/lib.rs b/crates/ml-core/src/lib.rs index a84faae6e..d45251f7c 100644 --- a/crates/ml-core/src/lib.rs +++ b/crates/ml-core/src/lib.rs @@ -145,7 +145,6 @@ pub mod nvtx; pub mod trading_action; pub mod action_space; pub mod xavier_init; -pub mod mixed_precision; pub mod order_router; pub mod fill_simulator; pub mod portfolio_tracker; diff --git a/crates/ml-core/src/mixed_precision.rs b/crates/ml-core/src/mixed_precision.rs deleted file mode 100644 index 4791cb8bb..000000000 --- a/crates/ml-core/src/mixed_precision.rs +++ /dev/null @@ -1,745 +0,0 @@ -//! Automatic Mixed Precision (AMP) utilities for 2x speedup -//! -//! This module provides utilities for mixed precision training: -//! - Forward pass in FP16/BF16 for 2x compute speedup -//! - Backward pass in FP32 for numerical stability -//! - Automatic dtype conversion between precision levels -//! -//! Wave 26 P2.1: AMP implementation for production DQN training - -use candle_core::{DType, Device, Tensor}; -use serde::{Deserialize, Serialize}; - -use crate::MLError; - -/// Round up a dimension to the next multiple of 8 for tensor core alignment. -/// -/// BF16/FP16 tensor cores (Volta+) require M, N, K dimensions to be multiples -/// of 8 for HMMA (half-precision matrix multiply-accumulate) instructions. -/// Unaligned dimensions cause cuBLAS to fall back to scalar FMA ops, leaving -/// tensor cores completely idle (0% utilization). -/// -/// Only pads on CUDA devices where tensor cores are available. -/// Returns the original dimension on CPU/Metal. -#[inline] -pub fn align_dim_for_tensor_cores(dim: usize, device: &Device) -> usize { - match device { - Device::Cuda(_) => (dim + 7) & !7, - Device::Cpu | Device::Metal(_) => dim, - } -} - -/// Configuration for mixed precision training -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MixedPrecisionConfig { - /// Whether mixed precision is enabled - pub enabled: bool, - /// Data type for reduced precision (F16 or BF16) - /// F16: Better on most GPUs, wider hardware support - /// BF16: Better range, preferred on modern hardware (Ampere+) - pub dtype: DTypeSelection, - /// Loss scaling factor to prevent gradient underflow - /// Typical values: 128.0 - 65536.0 - /// Set to 1.0 for BF16 (better range, less scaling needed) - pub loss_scale: f32, - /// Whether to dynamically adjust loss scale - pub dynamic_loss_scale: bool, - /// Growth factor for dynamic loss scaling - pub scale_growth_factor: f32, - /// Backoff factor for dynamic loss scaling - pub scale_backoff_factor: f32, - /// Number of consecutive steps without overflow before growing scale - pub scale_growth_interval: usize, -} - -/// Data type selection for mixed precision -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -pub enum DTypeSelection { - /// Half precision floating point (16-bit) - /// Range: ±65,504, Precision: ~3 decimal digits - /// Best for: Volta/Turing/Ampere GPUs - F16, - /// Brain floating point (16-bit) - /// Range: same as F32, Precision: ~2 decimal digits - /// Best for: Ampere+ GPUs with BF16 hardware support - BF16, -} - -impl DTypeSelection { - /// Convert to Candle DType - pub fn to_dtype(&self) -> DType { - match self { - DTypeSelection::F16 => DType::F16, - DTypeSelection::BF16 => DType::BF16, - } - } -} - -impl Default for MixedPrecisionConfig { - fn default() -> Self { - Self { - enabled: false, - dtype: DTypeSelection::F16, // F16 has wider hardware support - loss_scale: 1024.0, // Conservative starting point - dynamic_loss_scale: true, - scale_growth_factor: 2.0, - scale_backoff_factor: 0.5, - scale_growth_interval: 2000, - } - } -} - -impl MixedPrecisionConfig { - /// Create a configuration optimized for modern GPUs (Ampere+) - pub fn for_ampere() -> Self { - Self { - enabled: true, - dtype: DTypeSelection::BF16, // BF16 is optimal on Ampere+ - loss_scale: 1.0, // BF16 has better range, less scaling needed - dynamic_loss_scale: false, // BF16 rarely needs dynamic scaling - ..Default::default() - } - } - - /// Create FP16 configuration for Volta/Turing GPUs. - /// Alias for `for_volta_turing()` — a convenient default when you know - /// the GPU supports FP16 tensor cores but not necessarily BF16. - pub fn default_enabled() -> Self { - Self { - enabled: true, - dtype: DTypeSelection::F16, - loss_scale: 1024.0, - dynamic_loss_scale: true, - scale_growth_factor: 2.0, - scale_backoff_factor: 0.5, - scale_growth_interval: 2000, - } - } - - /// Create a configuration for older GPUs (Volta/Turing) - pub fn for_volta_turing() -> Self { - Self { - enabled: true, - dtype: DTypeSelection::F16, - loss_scale: 2048.0, // F16 needs more aggressive scaling - dynamic_loss_scale: true, - scale_growth_factor: 2.0, - scale_backoff_factor: 0.5, - scale_growth_interval: 2000, - } - } - - /// Create a disabled configuration (FP32 only) - pub fn disabled() -> Self { - Self { - enabled: false, - ..Default::default() - } - } -} - -/// Auto-detect mixed precision config from the GPU name returned by `nvidia-smi`. -/// -/// BF16 (Ampere+): A100, A10, A30, H100, H200, L4, L40, RTX 30xx, RTX 40xx, RTX 50xx -/// FP16 (Volta/Turing): V100, T4, RTX 20xx, Titan RTX, Quadro RTX -/// None: older GPUs, CPU, or detection failure -pub fn detect_from_gpu_name(gpu_name: &str) -> Option { - let name = gpu_name.to_uppercase(); - - // Ampere+ (compute capability >= 8.0) -- BF16 native support - let is_ampere_plus = name.contains("A100") - || name.contains("A10G") - || name.contains("A10 ") - || name.contains("A30") - || name.contains("A40") - || name.contains("A6000") - || name.contains("H100") - || name.contains("H200") - || name.contains("L4") - || name.contains("L40") - || name.contains("RTX 30") - || name.contains("RTX 40") - || name.contains("RTX 50") - || name.contains("RTX A") - || name.contains("3050") - || name.contains("3060") - || name.contains("3070") - || name.contains("3080") - || name.contains("3090") - || name.contains("4060") - || name.contains("4070") - || name.contains("4080") - || name.contains("4090") - || name.contains("5070") - || name.contains("5080") - || name.contains("5090"); - - if is_ampere_plus { - return Some(MixedPrecisionConfig::for_ampere()); - } - - // Volta/Turing (compute capability 7.x) -- FP16 tensor cores - let is_volta_turing = name.contains("V100") - || name.contains("T4") - || name.contains("RTX 20") - || name.contains("TITAN RTX") - || name.contains("QUADRO RTX") - || name.contains("2060") - || name.contains("2070") - || name.contains("2080"); - - if is_volta_turing { - return Some(MixedPrecisionConfig::for_volta_turing()); - } - - None -} - -/// Auto-detect mixed precision config from the GPU detected by `nvidia-smi`. -/// -/// Uses [`cached_capabilities()`](crate::gpu::capabilities::cached_capabilities) to -/// obtain the device name (run once, then cached) and forwards it to -/// [`detect_from_gpu_name`]. -/// -/// Returns `None` on CPU or if the GPU is too old for tensor-core mixed precision. -/// -/// Used by [`training_dtype`] to determine whether BF16 is available. -/// All CUDA devices use BF16 via `training_dtype`; this config is -/// used for logging and hardware capability reporting. -pub fn detect_from_gpu_name_auto() -> Option { - let caps = crate::gpu::capabilities::cached_capabilities(); - if !caps.is_cuda { - return None; - } - detect_from_gpu_name(&caps.device_name) -} - -/// Returns the training [`DType`] for the given device. -/// -/// - CUDA → `BF16` (tensor-core accelerated, same exponent range as F32). -/// - CPU / Metal → `F32`. -/// -/// This is the single source of truth for "what dtype should my weight -/// tensors and forward-pass intermediaries use?" Call it once when -/// constructing a model or trainer and thread the result through. -/// -/// **F32 at boundaries only**: Scalar tensors (rewards, dones, priorities) -/// and loss computation remain F32 for numerical stability. All state -/// tensors and model weights use BF16 on CUDA for full tensor-core -/// utilization. Use [`ensure_training_dtype`] at model entry points -/// to cast any stale F32 inputs. -pub fn training_dtype(device: &candle_core::Device) -> candle_core::DType { - match device { - candle_core::Device::Cuda(_) => candle_core::DType::BF16, - candle_core::Device::Cpu | candle_core::Device::Metal(_) => candle_core::DType::F32, - } -} - -/// Cast a tensor to the training dtype for its device if needed. -/// -/// This is the canonical "boundary cast" — call it at any model entry -/// point where the input tensor may arrive as F32 from external code -/// (tests, ensemble adapters, raw inference). When dtypes already -/// match the call is essentially free (Candle returns a clone). -pub fn ensure_training_dtype(tensor: &candle_core::Tensor) -> Result { - let target = training_dtype(tensor.device()); - if tensor.dtype() == target { - Ok(tensor.clone()) - } else { - tensor.to_dtype(target) - } -} - -/// Convert tensor to half precision (FP16 or BF16) -/// -/// # Arguments -/// * `tensor` - Input tensor (typically F32) -/// * `dtype` - Target half precision type -/// -/// # Returns -/// Tensor converted to specified half precision type -/// -/// # Errors -/// Returns MLError if conversion fails -pub fn to_half(tensor: &Tensor, dtype: DTypeSelection) -> Result { - tensor - .to_dtype(dtype.to_dtype()) - .map_err(|e| MLError::TensorOperationError(format!("Failed to convert to half precision: {}", e))) -} - -/// Convert tensor to full precision (F32) -/// -/// # Arguments -/// * `tensor` - Input tensor (typically F16/BF16) -/// -/// # Returns -/// Tensor converted to F32 -/// -/// # Errors -/// Returns MLError if conversion fails -pub fn to_float(tensor: &Tensor) -> Result { - tensor - .to_dtype(DType::F32) - .map_err(|e| MLError::TensorOperationError(format!("Failed to convert to float: {}", e))) -} - -/// Execute forward pass in mixed precision -/// -/// Pattern: Forward in FP16/BF16, backward in FP32 -/// This provides 2x speedup on forward pass while maintaining gradient stability -/// -/// # Arguments -/// * `input` - Input tensor in F32 -/// * `config` - Mixed precision configuration -/// * `forward_fn` - Forward pass function to execute -/// -/// # Returns -/// Output tensor in F32 (automatically converted back) -/// -/// # Errors -/// Returns MLError if conversion or forward pass fails -pub fn forward_mixed( - input: &Tensor, - config: &MixedPrecisionConfig, - forward_fn: F, -) -> Result -where - F: Fn(&Tensor) -> Result, -{ - if !config.enabled { - // Bypass: run in full precision - return forward_fn(input); - } - - // Convert input to half precision - let half_input = to_half(input, config.dtype)?; - - // Execute forward pass in half precision - let half_output = forward_fn(&half_input)?; - - // Convert output back to full precision for gradient computation - to_float(&half_output) -} - -/// Scale loss for mixed precision training -/// -/// Scales loss by a factor to prevent gradient underflow in FP16 -/// Gradients are automatically unscaled during optimizer step -/// -/// # Arguments -/// * `loss` - Loss tensor -/// * `scale` - Scaling factor -/// -/// # Returns -/// Scaled loss tensor -/// -/// # Errors -/// Returns MLError if scaling fails -pub fn scale_loss(loss: &Tensor, scale: f32) -> Result { - loss.affine(scale as f64, 0.0) - .map_err(|e| MLError::TensorOperationError(format!("Failed to scale loss: {}", e))) -} - -/// Unscale gradients after backward pass -/// -/// Divides gradients by scale factor to restore original magnitude -/// -/// # Arguments -/// * `gradients` - Gradient tensor -/// * `scale` - Scaling factor used for loss -/// -/// # Returns -/// Unscaled gradient tensor -/// -/// # Errors -/// Returns MLError if unscaling fails -pub fn unscale_gradients(gradients: &Tensor, scale: f32) -> Result { - gradients - .affine(1.0 / scale as f64, 0.0) - .map_err(|e| MLError::TensorOperationError(format!("Failed to unscale gradients: {}", e))) -} - -#[cfg(test)] -mod tests { - use super::*; - use candle_core::Device; - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_mixed_precision_config_default() { - let config = MixedPrecisionConfig::default(); - assert!(!config.enabled); - assert_eq!(config.dtype, DTypeSelection::F16); - assert_eq!(config.loss_scale, 1024.0); - assert!(config.dynamic_loss_scale); - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_mixed_precision_config_ampere() { - let config = MixedPrecisionConfig::for_ampere(); - assert!(config.enabled); - assert_eq!(config.dtype, DTypeSelection::BF16); - assert_eq!(config.loss_scale, 1.0); - assert!(!config.dynamic_loss_scale); - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_mixed_precision_config_volta_turing() { - let config = MixedPrecisionConfig::for_volta_turing(); - assert!(config.enabled); - assert_eq!(config.dtype, DTypeSelection::F16); - assert_eq!(config.loss_scale, 2048.0); - assert!(config.dynamic_loss_scale); - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_dtype_selection_to_dtype() { - assert_eq!(DTypeSelection::F16.to_dtype(), DType::F16); - assert_eq!(DTypeSelection::BF16.to_dtype(), DType::BF16); - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_to_half_f16() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let tensor = Tensor::ones((2, 3), DType::F32, &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let half_tensor = to_half(&tensor, DTypeSelection::F16)?; - - assert_eq!(half_tensor.dtype(), DType::F16); - assert_eq!(half_tensor.dims(), tensor.dims()); - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_to_half_bf16() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let tensor = Tensor::ones((2, 3), DType::F32, &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let half_tensor = to_half(&tensor, DTypeSelection::BF16)?; - - assert_eq!(half_tensor.dtype(), DType::BF16); - assert_eq!(half_tensor.dims(), tensor.dims()); - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_to_float() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let tensor = Tensor::ones((2, 3), DType::F16, &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let float_tensor = to_float(&tensor)?; - - assert_eq!(float_tensor.dtype(), DType::F32); - assert_eq!(float_tensor.dims(), tensor.dims()); - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_round_trip_conversion() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let original = Tensor::new(&[1.0_f32, 2.0, 3.0, 4.0], &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - // F32 -> F16 -> F32 - let half = to_half(&original, DTypeSelection::F16)?; - let restored = to_float(&half)?; - - assert_eq!(restored.dtype(), DType::F32); - assert_eq!(restored.dims(), original.dims()); - - // Values should be approximately equal (F16 has limited precision) - let original_data = original - .to_vec1::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - let restored_data = restored - .to_vec1::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - for (a, b) in original_data.iter().zip(restored_data.iter()) { - assert!((a - b).abs() < 1e-2, "Values differ: {} vs {}", a, b); - } - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_forward_mixed_disabled() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let input = Tensor::ones((2, 3), DType::F32, &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let config = MixedPrecisionConfig::disabled(); - - let forward_fn = |x: &Tensor| -> Result { - // Simple operation: multiply by 2 - x.affine(2.0, 0.0) - .map_err(|e| MLError::TensorOperationError(e.to_string())) - }; - - let output = forward_mixed(&input, &config, forward_fn)?; - - assert_eq!(output.dtype(), DType::F32); - // Use to_vec2 for 2D tensor - let data = output - .to_vec2::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - for row in &data { - for &val in row { - assert!((val - 2.0).abs() < 1e-6); - } - } - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_forward_mixed_enabled_f16() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let input = Tensor::ones((2, 3), DType::F32, &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let config = MixedPrecisionConfig::for_volta_turing(); - - let forward_fn = |x: &Tensor| -> Result { - // Verify input is F16 - assert_eq!(x.dtype(), DType::F16); - - // Simple operation: multiply by 2 - x.affine(2.0, 0.0) - .map_err(|e| MLError::TensorOperationError(e.to_string())) - }; - - let output = forward_mixed(&input, &config, forward_fn)?; - - // Output should be converted back to F32 - assert_eq!(output.dtype(), DType::F32); - // Use to_vec2 for 2D tensor - let data = output - .to_vec2::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - for row in &data { - for &val in row { - assert!((val - 2.0).abs() < 1e-2); // F16 has lower precision - } - } - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_forward_mixed_enabled_bf16() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let input = Tensor::ones((2, 3), DType::F32, &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let config = MixedPrecisionConfig::for_ampere(); - - let forward_fn = |x: &Tensor| -> Result { - // Verify input is BF16 - assert_eq!(x.dtype(), DType::BF16); - - // Simple operation: multiply by 2 - x.affine(2.0, 0.0) - .map_err(|e| MLError::TensorOperationError(e.to_string())) - }; - - let output = forward_mixed(&input, &config, forward_fn)?; - - // Output should be converted back to F32 - assert_eq!(output.dtype(), DType::F32); - // Use to_vec2 for 2D tensor - let data = output - .to_vec2::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - for row in &data { - for &val in row { - assert!((val - 2.0).abs() < 1e-2); // BF16 has lower precision - } - } - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_scale_loss() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let loss = Tensor::new(&[1.0_f32], &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let scaled = scale_loss(&loss, 1024.0)?; - let data = scaled - .to_vec1::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - assert!((data[0] - 1024.0).abs() < 1e-4); - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_unscale_gradients() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let gradients = Tensor::new(&[1024.0_f32], &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let unscaled = unscale_gradients(&gradients, 1024.0)?; - let data = unscaled - .to_vec1::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - assert!((data[0] - 1.0).abs() < 1e-4); - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_scale_unscale_round_trip() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let original_loss = Tensor::new(&[0.5_f32], &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let scale = 2048.0; - let scaled = scale_loss(&original_loss, scale)?; - let unscaled = unscale_gradients(&scaled, scale)?; - - let original_data = original_loss - .to_vec1::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - let unscaled_data = unscaled - .to_vec1::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - assert!((original_data[0] - unscaled_data[0]).abs() < 1e-4); - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_mixed_precision_preserves_shape() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let input = Tensor::ones((4, 8, 16), DType::F32, &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let config = MixedPrecisionConfig::for_ampere(); - - let forward_fn = |x: &Tensor| -> Result { - Ok(x.clone()) - }; - - let output = forward_mixed(&input, &config, forward_fn)?; - - assert_eq!(output.dims(), input.dims()); - assert_eq!(output.dtype(), DType::F32); - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_mixed_precision_numerical_accuracy() -> Result<(), MLError> { - let device = Device::new_cuda(0).expect("CUDA required"); - let input = Tensor::new(&[1.5_f32, 2.5, 3.5, 4.5], &device) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - let config = MixedPrecisionConfig::for_volta_turing(); - - let forward_fn = |x: &Tensor| -> Result { - // Complex operation: (x * 2 + 1) / 3 - let mul = x - .affine(2.0, 0.0) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - let add = mul - .affine(1.0, 1.0) - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - add.affine(1.0 / 3.0, 0.0) - .map_err(|e| MLError::TensorOperationError(e.to_string())) - }; - - let output = forward_mixed(&input, &config, forward_fn)?; - let data = output - .to_vec1::() - .map_err(|e| MLError::TensorOperationError(e.to_string()))?; - - // Expected: (x * 2 + 1) / 3 for each input - let expected = [1.333333, 2.0, 2.666667, 3.333333]; - for (actual, &exp) in data.iter().zip(expected.iter()) { - assert!( - (actual - exp).abs() < 1e-2, - "Numerical accuracy error: {} vs {}", - actual, - exp - ); - } - - Ok(()) - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_detect_from_gpu_name_ampere() { - // Ampere+ GPUs should get BF16 config - for name in &[ - "NVIDIA A100-SXM4-80GB", - "NVIDIA H100 80GB HBM3", - "NVIDIA L4", - "NVIDIA GeForce RTX 3050 Ti Laptop GPU", - "NVIDIA GeForce RTX 3090", - "NVIDIA GeForce RTX 4090", - "NVIDIA RTX A6000", - ] { - let config = detect_from_gpu_name(name); - assert!(config.is_some(), "Expected BF16 config for {}", name); - let c = config.as_ref().unwrap(); - assert!(c.enabled, "Expected enabled for {}", name); - assert_eq!(c.dtype, DTypeSelection::BF16, "Expected BF16 for {}", name); - } - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_detect_from_gpu_name_volta_turing() { - // Volta/Turing GPUs should get FP16 config - for name in &[ - "Tesla V100-SXM2-16GB", - "Tesla T4", - "NVIDIA GeForce RTX 2080 Ti", - "Quadro RTX 8000", - ] { - let config = detect_from_gpu_name(name); - assert!(config.is_some(), "Expected FP16 config for {}", name); - let c = config.as_ref().unwrap(); - assert!(c.enabled, "Expected enabled for {}", name); - assert_eq!(c.dtype, DTypeSelection::F16, "Expected FP16 for {}", name); - } - } - - #[cfg_attr(not(feature = "cuda"), ignore)] - #[test] - fn test_detect_from_gpu_name_older() { - // Older GPUs or unknown strings should return None - for name in &["Tesla K80", "Tesla P100", "GeForce GTX 1080 Ti", "CPU", ""] { - let config = detect_from_gpu_name(name); - assert!(config.is_none(), "Expected None for {:?}", name); - } - } -} diff --git a/crates/ml-dqn/src/agent.rs b/crates/ml-dqn/src/agent.rs index f62f1e7fd..01d6a2473 100644 --- a/crates/ml-dqn/src/agent.rs +++ b/crates/ml-dqn/src/agent.rs @@ -8,7 +8,6 @@ use std::collections::HashMap; use ml_core::optimizers::Adam; use candle_core::Tensor; use candle_nn::{ops::leaky_relu, Module, VarBuilder}; -use crate::mixed_precision::training_dtype; use candle_optimisers::adam::ParamsAdam; // Use our Adam wrapper from lib.rs use serde::{Deserialize, Serialize}; use tracing::debug; @@ -207,7 +206,7 @@ impl DQNAgent { use_spectral_norm: false, spectral_norm_iterations: 1, use_residual: false, - mixed_precision: config.mixed_precision.clone(), + }; // Create Q-networks @@ -340,12 +339,12 @@ impl DQNAgent { // Forward pass through main network with gradient tracking let var_builder = - VarBuilder::from_varmap(self.q_network.vars(), training_dtype(device), device); + VarBuilder::from_varmap(self.q_network.vars(), candle_core::DType::BF16, device); let current_q_values = self.forward_with_gradients(&state_tensor, &var_builder)?; // Forward pass through target network WITHOUT gradients let target_var_builder = - VarBuilder::from_varmap(self.target_network.vars(), training_dtype(device), device); + VarBuilder::from_varmap(self.target_network.vars(), candle_core::DType::BF16, device); let next_q_values = self.forward_without_gradients(&next_state_tensor, &target_var_builder)?; @@ -364,7 +363,7 @@ impl DQNAgent { let max_next_q = next_q_values.max(1)?; // Get maximum values // Create reward and done tensors, cast to training dtype at the boundary - let dtype = training_dtype(device); + let dtype = candle_core::DType::BF16; let reward_tensor = Tensor::from_vec(rewards.to_vec(), batch_size, device).map_err(|e| { MLError::TrainingError(format!("Failed to create reward tensor: {}", e)) @@ -419,17 +418,9 @@ impl DQNAgent { ) -> Result { use candle_nn::linear; - // Mixed precision: cast input to reduced precision for 2x compute speedup - let (x_input, use_amp) = match &self.config.mixed_precision { - Some(mp) if mp.enabled => { - let target_dtype = mp.dtype.to_dtype(); - match input.to_dtype(target_dtype) { - Ok(converted) => (converted, true), - Err(_) => (input.clone(), false), - } - } - _ => (input.clone(), false), - }; + // BF16 on CUDA, F32 on CPU + let x_input = input.to_dtype(candle_core::DType::BF16) + .map_err(|e| MLError::ModelError(format!("dtype cast: {e}")))?; let mut layers = Vec::new(); let mut input_dim = self.config.state_dim; @@ -467,18 +458,13 @@ impl DQNAgent { } } - // Mixed precision: cast output back to FP32 for loss computation - if use_amp { - x = x.to_dtype(candle_core::DType::F32)?; - } + // F32 at boundary + x = x.to_dtype(candle_core::DType::F32)?; Ok(x) } /// Forward pass through network without gradient tracking (for target network) - /// - /// Supports mixed precision: casts input to BF16/FP16 for compute, - /// casts output back to FP32 for target Q-value calculation. fn forward_without_gradients( &self, input: &Tensor, @@ -486,17 +472,9 @@ impl DQNAgent { ) -> Result { use candle_nn::linear; - // Mixed precision: cast input to reduced precision - let (x_input, use_amp) = match &self.config.mixed_precision { - Some(mp) if mp.enabled => { - let target_dtype = mp.dtype.to_dtype(); - match input.to_dtype(target_dtype) { - Ok(converted) => (converted, true), - Err(_) => (input.clone(), false), - } - } - _ => (input.clone(), false), - }; + // BF16 on CUDA, F32 on CPU + let x_input = input.to_dtype(candle_core::DType::BF16) + .map_err(|e| MLError::ModelError(format!("dtype cast: {e}")))?; let mut layers = Vec::new(); let mut input_dim = self.config.state_dim; @@ -530,13 +508,8 @@ impl DQNAgent { } } - // Mixed precision: cast output back to FP32 - if use_amp { - x = x.to_dtype(candle_core::DType::F32)?; - } - - // Detach from gradient computation - Ok(x.detach()) + // F32 at boundary, detach from gradient computation + Ok(x.to_dtype(candle_core::DType::F32)?.detach()) } fn update_target_network_weights(&mut self) -> Result<(), MLError> { diff --git a/crates/ml-dqn/src/attention.rs b/crates/ml-dqn/src/attention.rs index 51d60a337..68359a0c8 100644 --- a/crates/ml-dqn/src/attention.rs +++ b/crates/ml-dqn/src/attention.rs @@ -255,7 +255,7 @@ impl MultiHeadAttention { /// 5. Optional: Add residual connection and layer normalization pub fn forward(&self, x: &Tensor, mask: Option<&Tensor>) -> Result { // Cast input to training dtype (BF16 on CUDA, F32 on CPU) - let x = &crate::mixed_precision::ensure_training_dtype(x) + let x = &x.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input dtype: {}", e)))?; let residual = x.clone(); @@ -433,7 +433,6 @@ impl MultiHeadAttention { mod tests { use super::*; use candle_nn::VarMap; - use crate::mixed_precision::training_dtype; #[test] fn test_config_validation() { @@ -471,7 +470,7 @@ mod tests { let device = Device::new_cuda(0).expect("CUDA required"); let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; assert_eq!(attention.config().embed_dim, 64); @@ -485,7 +484,7 @@ mod tests { let device = Device::new_cuda(0).expect("CUDA required"); let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -515,7 +514,7 @@ mod tests { let device = Device::new_cuda(0).expect("CUDA required"); let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -555,7 +554,7 @@ mod tests { let device = Device::new_cuda(0).expect("CUDA required"); let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -589,7 +588,7 @@ mod tests { config.use_layer_norm = false; // Disable to test residual alone let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -620,7 +619,7 @@ mod tests { let embed_dim = 64; let config = MultiHeadAttentionConfig::new(embed_dim, num_heads)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; diff --git a/crates/ml-dqn/src/branching.rs b/crates/ml-dqn/src/branching.rs index ae3d52d86..dd360220b 100644 --- a/crates/ml-dqn/src/branching.rs +++ b/crates/ml-dqn/src/branching.rs @@ -40,10 +40,8 @@ use candle_core::{DType, Device, ModuleT, Tensor, Var}; use candle_nn::{Dropout, Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; -use crate::mixed_precision::training_dtype; use crate::noisy_layers::NoisyLinear; use crate::xavier_init::linear_xavier; -use ml_core::mixed_precision::align_dim_for_tensor_cores; use ml_core::MLError; /// Output of the branching network's forward pass. @@ -173,7 +171,7 @@ impl BranchingConfig { branch_sizes: Vec, ) -> Self { let aligned_state_dim = match device { - Some(d) => align_dim_for_tensor_cores(state_dim, d), + Some(_d) => (state_dim + 7) & !7, None => state_dim, }; Self { @@ -327,7 +325,7 @@ impl BranchingDuelingQNetwork { } let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); // Shared encoder (always standard Linear -- noise only in heads) let mut shared_layers = Vec::new(); @@ -529,7 +527,7 @@ impl BranchingDuelingQNetwork { train: bool, ) -> Result { // Shared encoder - let mut h = crate::mixed_precision::ensure_training_dtype(state) + let mut h = state.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; for (i, layer) in self.shared_layers.iter().enumerate() { @@ -1917,14 +1915,14 @@ mod tests { #[test] fn test_alignment_function_directly() { // On CUDA, align to next multiple of 8 for tensor core HMMA dispatch - let aligned_gpu = align_dim_for_tensor_cores(45, &cuda_device()); + let aligned_gpu = (45 + 7) & !7; assert_eq!(aligned_gpu, 48); // 45 → 48 (next multiple of 8) // Already aligned values remain unchanged - let aligned_gpu_8 = align_dim_for_tensor_cores(48, &cuda_device()); + let aligned_gpu_8 = (48 + 7) & !7; assert_eq!(aligned_gpu_8, 48); - let aligned_gpu_1 = align_dim_for_tensor_cores(1, &cuda_device()); + let aligned_gpu_1 = (1 + 7) & !7; assert_eq!(aligned_gpu_1, 8); // 1 → 8 (next multiple of 8) } diff --git a/crates/ml-dqn/src/curiosity.rs b/crates/ml-dqn/src/curiosity.rs index cf76b1806..440e52219 100644 --- a/crates/ml-dqn/src/curiosity.rs +++ b/crates/ml-dqn/src/curiosity.rs @@ -7,7 +7,6 @@ use candle_core::{Device, Tensor}; use candle_nn::{ops::leaky_relu, AdamW, Linear, Module, Optimizer, ParamsAdamW, VarBuilder, VarMap}; use super::action_space::{FactoredAction, ExposureLevel}; -use super::mixed_precision::training_dtype; use ml_core::MLError; use crate::xavier_init::linear_xavier; @@ -49,7 +48,7 @@ impl ForwardDynamicsModel { action_categories: usize, ) -> Result { let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let input_dim = market_dim + action_categories; let fc1 = linear_xavier(input_dim, hidden_dim, var_builder.pp("fc1")) @@ -83,12 +82,12 @@ impl ForwardDynamicsModel { // Extract first market_dim features from state (market features only, skip portfolio/OFI) let state_embedding = state.narrow(1, 0, self.market_dim) .map_err(|e| MLError::ModelError(format!("Failed to narrow state: {}", e)))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to convert state dtype: {}", e)))?; // One-hot encode action (convert FactoredAction to simplified action index) let batch_size = state.dims()[0]; - let mut action_onehot = Tensor::zeros((batch_size, self.action_categories), training_dtype(&self.device), &self.device) + let mut action_onehot = Tensor::zeros((batch_size, self.action_categories), candle_core::DType::BF16, &self.device) .map_err(|e| MLError::ModelError(format!("Failed to create action tensor: {}", e)))?; // Convert FactoredAction to simplified category: 0=SHORT, 1=FLAT, 2=LONG @@ -98,7 +97,7 @@ impl ForwardDynamicsModel { ExposureLevel::Long50 | ExposureLevel::Long100 => 2_i64, // LONG }; for batch_idx in 0..batch_size { - action_onehot = action_onehot.slice_assign(&[batch_idx..batch_idx+1, action_idx as usize..action_idx as usize+1], &Tensor::ones((1, 1), training_dtype(&self.device), &self.device)?) + action_onehot = action_onehot.slice_assign(&[batch_idx..batch_idx+1, action_idx as usize..action_idx as usize+1], &Tensor::ones((1, 1), candle_core::DType::BF16, &self.device)?) .map_err(|e| MLError::ModelError(format!("Failed to set action one-hot: {}", e)))?; } @@ -226,7 +225,7 @@ impl CuriosityModule { let market_dim = self.forward_model.market_dim; let next_state_embedding = next_state.narrow(1, 0, market_dim) .map_err(|e| MLError::ModelError(format!("Failed to narrow next_state: {}", e)))? - .to_dtype(training_dtype(&self.forward_model.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to convert next_state dtype: {}", e)))?; // Predict next state @@ -305,7 +304,7 @@ mod tests { // Create state and target (cast to training dtype for BF16 compat) let state = Tensor::randn(0.0_f32, 1.0, (4, MARKET_DIM), &device)?; let target = Tensor::randn(0.0_f32, 1.0, (4, MARKET_DIM), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; let action = test_buy_action(); // Get initial prediction (BF16 output) diff --git a/crates/ml-dqn/src/distributional_dueling.rs b/crates/ml-dqn/src/distributional_dueling.rs index 6658372e8..bb9f06d3e 100644 --- a/crates/ml-dqn/src/distributional_dueling.rs +++ b/crates/ml-dqn/src/distributional_dueling.rs @@ -43,7 +43,6 @@ use candle_core::{Device, Tensor}; use candle_nn::{Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; -use crate::mixed_precision::training_dtype; use crate::rmsnorm::RMSNorm; use crate::xavier_init::{linear_near_zero_init, linear_xavier}; use ml_core::MLError; @@ -161,7 +160,7 @@ impl DistributionalDuelingQNetwork { pub fn new(config: DistributionalDuelingConfig, device: Device) -> Result { // state_dim is pre-aligned to 8 by the caller for tensor core utilization let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); // Build shared feature layers with RMSNorm after each let mut shared_layers = Vec::new(); @@ -253,7 +252,7 @@ impl DistributionalDuelingQNetwork { /// - `A(s,a,z_i)`: Advantage distribution per action (probability of atom `z_i`) /// - `mean(A(s,·,z_i))`: Mean advantage across actions (ensures identifiability) pub fn forward(&self, state: &Tensor) -> Result { - let state = crate::mixed_precision::ensure_training_dtype(state) + let state = state.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let batch_size = state .dim(0) diff --git a/crates/ml-dqn/src/dqn.rs b/crates/ml-dqn/src/dqn.rs index 399a12c0c..a4e12658d 100644 --- a/crates/ml-dqn/src/dqn.rs +++ b/crates/ml-dqn/src/dqn.rs @@ -11,7 +11,6 @@ use std::collections::VecDeque; use std::sync::{Arc, Mutex}; -use crate::mixed_precision::training_dtype; use crate::target_update::{convergence_half_life, hard_update, polyak_update}; // WAVE 16 (Agent 36) use crate::xavier_init::linear_xavier; // Xavier initialization with VarMap registration use ml_core::optimizers::Adam; @@ -98,9 +97,6 @@ pub struct DQNConfig { /// Whether to use Prioritized Experience Replay pub use_per: bool, /// Whether to use GPU-resident replay buffer (only when `use_per=true` + CUDA). - /// When false, PER uses CPU-side buffer even on CUDA devices. - /// Set to false on GPUs with ≤8 GB VRAM to avoid dual-allocator fragmentation. - pub use_gpu_replay_buffer: bool, /// PER alpha parameter (prioritization exponent) pub per_alpha: f64, /// PER beta start value (importance sampling weight) @@ -109,10 +105,10 @@ pub struct DQNConfig { pub per_beta_max: f64, /// Number of steps to anneal beta from start to max pub per_beta_annealing_steps: usize, - /// Maximum memory allocation for GPU/CPU PER replay buffer (bytes). + /// Maximum memory allocation for GPU PER replay buffer (bytes). /// Computed at runtime from detected GPU VRAM via /// `GpuHardwareInfo::per_max_buffer_bytes()`. - /// Default: 4 GB (for tests/CPU). Overridden by `DqnTrainer` on CUDA. + /// Overridden by `DqnTrainer` constructor on CUDA. pub per_max_memory_bytes: usize, // Dueling Networks configuration @@ -262,9 +258,6 @@ pub struct DQNConfig { /// Applied after each shared layer activation during training forward pass only. pub dropout_rate: f64, - /// Mixed precision configuration for BF16/FP16 forward pass on supported GPUs. - /// None = FP32 only. Auto-configured based on GPU architecture at runtime. - pub mixed_precision: Option, } impl Default for DQNConfig { @@ -296,7 +289,7 @@ impl Default for DQNConfig { n_steps: 1, initial_capital: 100_000.0, use_per: true, - use_gpu_replay_buffer: true, + per_alpha: 0.6, per_beta_start: 0.4, per_beta_max: 1.0, @@ -359,7 +352,7 @@ impl Default for DQNConfig { dropout_rate: 0.0, // Disabled by default; hyperopt wires from dropout_initial // Mixed precision: disabled by default (auto-detected at runtime) - mixed_precision: None, + } } } @@ -615,7 +608,7 @@ impl DQNConfig { n_steps: 3, initial_capital: 100000.0, use_per: true, - use_gpu_replay_buffer: true, + per_alpha: 0.6, per_beta_start: 0.4, per_beta_max: 1.0, @@ -650,7 +643,7 @@ impl DQNConfig { minimum_profit_factor: 1.5, weight_decay: 1e-4, - mixed_precision: None, + ..Default::default() } } @@ -685,7 +678,7 @@ impl DQNConfig { n_steps: 1, // Default to single-step TD (most stable) initial_capital: 100000.0, use_per: true, // GPU PER mandatory on CUDA - use_gpu_replay_buffer: true, + per_alpha: 0.6, per_beta_start: 0.4, per_beta_max: 1.0, @@ -720,7 +713,7 @@ impl DQNConfig { minimum_profit_factor: 2.0, // Higher safety margin for conservative config weight_decay: 1e-4, - mixed_precision: None, + ..Default::default() } } @@ -765,7 +758,7 @@ impl DQNConfig { // Rainbow DQN defaults (emergency mode still needs GPU PER on CUDA) initial_capital: 100000.0, use_per: true, - use_gpu_replay_buffer: true, + per_alpha: 0.6, per_beta_start: 0.4, per_beta_max: 1.0, @@ -817,7 +810,7 @@ impl DQNConfig { minimum_profit_factor: 1.5, weight_decay: 1e-4, dropout_rate: 0.0, - mixed_precision: None, + } } } @@ -1004,7 +997,7 @@ impl Sequential { // input_dim is expected to be pre-aligned to 8 (tensor core requirement) // by the caller (DQNConfig.state_dim is aligned at construction time). let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let mut layers = Vec::new(); let mut noisy_layers = Vec::new(); @@ -1060,7 +1053,7 @@ impl Sequential { /// Forward pass through network pub fn forward(&self, input: &Tensor) -> Result { - let mut x = crate::mixed_precision::ensure_training_dtype(input) + let mut x = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; if self.use_noisy_nets { @@ -1371,13 +1364,14 @@ impl DQN { // Create experience replay buffer (uniform or prioritized based on config) let memory = if config.use_per { - // GPU PER: on CUDA, allocate GPU-resident ring buffer with OOM fallback to CPU PER. + // GPU PER: on CUDA, allocate GPU-resident ring buffer (hard error on failure). // This activates the GpuBatch fast path in compute_loss_internal() and // GPU TD error retention — eliminating 5 of 6 CPU↔GPU roundtrips per train step. #[cfg(feature = "cuda")] { - if device.is_cuda() && config.use_gpu_replay_buffer { - super::replay_buffer_type::ReplayBufferType::try_gpu_prioritized_with_fallback( + if device.is_cuda() { + // GPU PER is mandatory on CUDA. + super::replay_buffer_type::ReplayBufferType::try_gpu_with_halving( config.replay_buffer_capacity, config.state_dim, config.per_alpha, @@ -1388,9 +1382,6 @@ impl DQN { &device, )? } else { - if device.is_cuda() && !config.use_gpu_replay_buffer { - tracing::info!("GPU PER disabled (use_gpu_replay_buffer=false), using CPU PER"); - } super::replay_buffer_type::ReplayBufferType::new_prioritized( config.replay_buffer_capacity, config.per_alpha, @@ -1568,7 +1559,7 @@ impl DQN { /// 3. Standard Q-network (fallback) pub fn forward(&self, state: &Tensor) -> Result { // Auto-convert input to correct device and dtype - let state = crate::mixed_precision::ensure_training_dtype(state) + let state = state.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let state = state .to_device(&self.device) @@ -1591,7 +1582,7 @@ impl DQN { // GPU-native: arange + affine creates atoms on device (no CPU Vec) // atoms[i] = v_min + i * delta_z - let dtype = training_dtype(&self.device); + let dtype = candle_core::DType::BF16; let atoms_tensor = Tensor::arange(0_u32, num_atoms as u32, &self.device)? .to_dtype(DType::F32)? .affine(delta_z as f64, v_min as f64)? @@ -2249,7 +2240,7 @@ impl DQN { MLError::ModelError("Branching enabled but network not initialized".into()) })?; // Ensure input is on the correct device and dtype (matches forward() contract) - let states = crate::mixed_precision::ensure_training_dtype(states) + let states = states.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let states = states.to_device(&self.device) .map_err(|e| MLError::ModelError(format!("device migration: {e}")))?; @@ -2434,7 +2425,7 @@ impl DQN { /// the intermediate representation needed by IQN. fn get_state_embedding(&self, states: &Tensor) -> Result { let states = states.to_device(&self.device)?; - let states = crate::mixed_precision::ensure_training_dtype(&states) + let states = states.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; if self.q_network.use_noisy_nets { @@ -2487,9 +2478,11 @@ impl DQN { #[cfg(feature = "cuda")] let gpu_batch_opt = batch_sample.gpu_batch; - // CPU-side experience/weight Vecs: needed for non-CUDA builds and as fallback - // when GPU PER isn't available (small test buffers on CUDA). + // CPU-side experience/weight Vecs: only needed for non-CUDA tensor construction. + // In the CUDA path, gpu_batch tensors replace these entirely. + #[cfg(not(feature = "cuda"))] let experiences = batch_sample.experiences; + #[cfg(not(feature = "cuda"))] let weights = batch_sample.weights; let indices = batch_sample.indices; @@ -2574,10 +2567,10 @@ impl DQN { let max_action = (effective_actions.saturating_sub(1)) as f64; break 'tensor_prep ( bs, - gpu.states.to_dtype(training_dtype(device)).map_err(|e| { + gpu.states.to_dtype(candle_core::DType::BF16).map_err(|e| { MLError::TrainingError(format!("GPU states dtype cast: {}", e)) })?, - gpu.next_states.to_dtype(training_dtype(device)).map_err(|e| { + gpu.next_states.to_dtype(candle_core::DType::BF16).map_err(|e| { MLError::TrainingError(format!("GPU next_states dtype cast: {}", e)) })?, gpu.actions.clamp(0.0_f64, max_action).map_err(|e| { @@ -2596,9 +2589,14 @@ impl DQN { } } - // CPU PER fallback: build tensors from experience data via Tensor::new. - // On CUDA this uploads to GPU per-batch (acceptable for small test buffers - // where GPU PER couldn't allocate). Production always uses GPU PER. + // CUDA production: GpuBatch is mandatory — no CPU path. + #[cfg(feature = "cuda")] + return Err(MLError::TrainingError( + "GPU batch required — GPU PER is mandatory for DQN training.".to_owned(), + )); + + // Non-CUDA (tests only): build tensors from experience data via Tensor::new. + #[cfg(not(feature = "cuda"))] { let batch_size = experiences.len(); let state_dim = self.config.state_dim; @@ -2628,11 +2626,11 @@ impl DQN { let states_tensor = Tensor::new(states.as_slice(), device) .and_then(|t| t.reshape((batch_size, state_dim))) - .and_then(|t| t.to_dtype(training_dtype(device))) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| MLError::TrainingError(format!("States tensor: {}", e)))?; let next_states_tensor = Tensor::new(next_states.as_slice(), device) .and_then(|t| t.reshape((batch_size, state_dim))) - .and_then(|t| t.to_dtype(training_dtype(device))) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| MLError::TrainingError(format!("Next states tensor: {}", e)))?; let actions_tensor = Tensor::new(actions.as_slice(), device) .map_err(|e| MLError::TrainingError(format!("Actions tensor: {}", e)))?; @@ -2961,7 +2959,11 @@ impl DQN { #[cfg(not(feature = "cuda"))] return Err(MLError::TrainingError("GPU PER requires cuda feature".to_owned())) } else { - // CPU PER fallback: download TD errors from GPU for priority update. + #[cfg(feature = "cuda")] + return Err(MLError::TrainingError( + "CPU PER fallback disabled — use ReplayBufferType::GpuPrioritized when cuda is enabled".to_owned() + )); + #[cfg(not(feature = "cuda"))] (td_errors_for_per.to_dtype(DType::F32)?.to_vec1()?, None::, None::) } } else { @@ -3429,7 +3431,11 @@ impl DQN { #[cfg(not(feature = "cuda"))] return Err(MLError::TrainingError("GPU PER requires cuda feature".to_owned())) } else { - // CPU PER fallback: download TD errors from GPU for priority update. + #[cfg(feature = "cuda")] + return Err(MLError::TrainingError( + "CPU PER fallback disabled — use ReplayBufferType::GpuPrioritized when cuda is enabled".to_owned() + )); + #[cfg(not(feature = "cuda"))] (diff.detach().to_dtype(DType::F32)?.to_vec1()?, None::, None::) } } else { diff --git a/crates/ml-dqn/src/dueling.rs b/crates/ml-dqn/src/dueling.rs index bbb904424..6f8e46dfd 100644 --- a/crates/ml-dqn/src/dueling.rs +++ b/crates/ml-dqn/src/dueling.rs @@ -34,7 +34,6 @@ use candle_core::{Device, ModuleT, Tensor}; use candle_nn::{Dropout, Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; -use crate::mixed_precision::training_dtype; use crate::xavier_init::linear_xavier; use ml_core::MLError; @@ -144,7 +143,7 @@ impl DuelingQNetwork { pub fn new(config: DuelingConfig, device: Device) -> Result { // state_dim is pre-aligned to 8 by the caller for tensor core utilization let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); // Build shared feature layers let mut shared_layers = Vec::new(); @@ -226,7 +225,7 @@ impl DuelingQNetwork { /// When `train=true`, dropout is applied after each shared layer activation /// to regularize and prevent overfitting. pub fn forward_t(&self, state: &Tensor, train: bool) -> Result { - let mut h = crate::mixed_precision::ensure_training_dtype(state) + let mut h = state.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; for (i, layer) in self.shared_layers.iter().enumerate() { h = layer.forward(&h).map_err(|e| { diff --git a/crates/ml-dqn/src/gpu_replay_buffer.rs b/crates/ml-dqn/src/gpu_replay_buffer.rs index c1b953c84..fc73b686a 100644 --- a/crates/ml-dqn/src/gpu_replay_buffer.rs +++ b/crates/ml-dqn/src/gpu_replay_buffer.rs @@ -9,7 +9,6 @@ use candle_core::{CpuStorage, Device, DType, Layout, Shape, Tensor}; use ml_core::nvtx::NvtxRange; use ml_core::MLError; -use crate::mixed_precision::training_dtype; use crate::replay_buffer_type::GpuBatch; // --------------------------------------------------------------------------- @@ -356,7 +355,7 @@ impl GpuReplayBuffer { ))); } - let dtype = training_dtype(device); + let dtype = candle_core::DType::BF16; let states = Tensor::zeros(&[cap, sdim], dtype, device)?; let next_states = Tensor::zeros(&[cap, sdim], dtype, device)?; let actions = Tensor::zeros(&[cap], DType::U32, device)?; diff --git a/crates/ml-dqn/src/lib.rs b/crates/ml-dqn/src/lib.rs index 35d8d11b6..b06b93df7 100644 --- a/crates/ml-dqn/src/lib.rs +++ b/crates/ml-dqn/src/lib.rs @@ -18,7 +18,6 @@ // Re-export shared modules from ml-core for convenience pub use ml_core::action_space; pub use ml_core::order_router; -pub use ml_core::mixed_precision; pub use ml_core::xavier_init; pub use ml_core::portfolio_tracker; pub use ml_core::trading_action; diff --git a/crates/ml-dqn/src/network.rs b/crates/ml-dqn/src/network.rs index 06e8290c7..d77c8c470 100644 --- a/crates/ml-dqn/src/network.rs +++ b/crates/ml-dqn/src/network.rs @@ -3,7 +3,6 @@ use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering}; use candle_core::{Device, Result as CandleResult, Tensor}; -use crate::mixed_precision::training_dtype; use candle_nn::Module; use candle_nn::{ops::leaky_relu, Dropout, Linear, VarBuilder, VarMap}; use rand::prelude::*; // Replace common::rng with standard rand @@ -91,9 +90,6 @@ pub struct QNetworkConfig { pub spectral_norm_iterations: usize, /// Whether to use residual connections (Wave 26 P0.4) pub use_residual: bool, - /// Mixed precision configuration for tensor core utilization. - /// None = FP32 only. Some(config) = BF16 or FP16 forward pass. - pub mixed_precision: Option, } impl Default for QNetworkConfig { @@ -113,7 +109,7 @@ impl Default for QNetworkConfig { use_spectral_norm: false, // Disabled by default, enable for unstable training spectral_norm_iterations: 1, // 1 iteration is typically sufficient use_residual: false, // Disabled by default, opt-in for deeper networks (Wave 26 P0.4) - mixed_precision: None, // FP32 only by default; auto-configured when GPU detected + } } } @@ -189,32 +185,11 @@ impl NetworkLayers { Ok(Self { layers, dropout, training }) } - /// Forward pass with mixed precision: compute in BF16 on CUDA, return F32. - fn forward_mixed( - &self, - xs: &Tensor, - _mixed_precision: &Option, - ) -> CandleResult { - let xs = crate::mixed_precision::ensure_training_dtype(xs)?; - - let mut x = xs; - - for (i, layer) in self.layers.iter().enumerate() { - x = layer.forward(&x)?; - if i < self.layers.len() - 1 { - x = leaky_relu(&x, 0.01)?; - x = self.dropout.forward(&x, self.training)?; - } - } - - // F32 at boundary: downstream code (softmax, loss, value extraction) expects F32 - x.to_dtype(candle_core::DType::F32) - } } impl Module for NetworkLayers { fn forward(&self, xs: &Tensor) -> CandleResult { - let mut x = crate::mixed_precision::ensure_training_dtype(xs)?; + let mut x = xs.to_dtype(candle_core::DType::BF16)?; // Forward through hidden layers with LeakyReLU activation and dropout // LeakyReLU prevents dead neurons (0.01 gradient for negative inputs vs 0 for ReLU) @@ -247,12 +222,12 @@ impl QNetwork { let target_vars = VarMap::new(); // Initialize network weights - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let _layers = NetworkLayers::new(&var_builder, &config, &device, false) .map_err(|e| MLError::ModelError(format!("Failed to create network layers: {}", e)))?; // Initialize target network with same architecture - let target_var_builder = VarBuilder::from_varmap(&target_vars, training_dtype(&device), &device); + let target_var_builder = VarBuilder::from_varmap(&target_vars, candle_core::DType::BF16, &device); let _target_layers = NetworkLayers::new(&target_var_builder, &config, &device, false).map_err(|e| { MLError::ModelError(format!("Failed to create target network layers: {}", e)) @@ -293,7 +268,7 @@ impl QNetwork { // Get current dropout rate (adaptive or static) let dropout_rate = self.get_dropout_rate(); - let var_builder = VarBuilder::from_varmap(&self.vars, training_dtype(&self.device), &self.device); + let var_builder = VarBuilder::from_varmap(&self.vars, candle_core::DType::BF16, &self.device); let layers = NetworkLayers::new_with_dropout_rate( &var_builder, &self.config, @@ -309,7 +284,7 @@ impl QNetwork { .map_err(|e| MLError::ModelError(format!("Failed to add batch dimension: {}", e)))?; let output = layers - .forward_mixed(&input, &self.config.mixed_precision) + .forward(&input) .map_err(|e| MLError::ModelError(format!("Forward pass failed: {}", e)))?; let output_vec = output @@ -358,7 +333,7 @@ impl QNetwork { flat_states.extend_from_slice(state); } - let var_builder = VarBuilder::from_varmap(&self.vars, training_dtype(&self.device), &self.device); + let var_builder = VarBuilder::from_varmap(&self.vars, candle_core::DType::BF16, &self.device); let layers = NetworkLayers::new(&var_builder, &self.config, &self.device, self.is_training()) .map_err(|e| MLError::ModelError(format!("Failed to create layers: {}", e)))?; @@ -366,7 +341,7 @@ impl QNetwork { .map_err(|e| MLError::ModelError(format!("Failed to create input tensor: {}", e)))?; let output = layers - .forward_mixed(&input, &self.config.mixed_precision) + .forward(&input) .map_err(|e| MLError::ModelError(format!("Forward pass failed: {}", e)))?; let output_vec = output.to_vec2::().map_err(|e| { diff --git a/crates/ml-dqn/src/noisy_layers.rs b/crates/ml-dqn/src/noisy_layers.rs index 66e7010e1..54d889dcd 100644 --- a/crates/ml-dqn/src/noisy_layers.rs +++ b/crates/ml-dqn/src/noisy_layers.rs @@ -13,7 +13,6 @@ use candle_core::{Device, Result as CandleResult, Tensor, Var}; use candle_nn::{Module, VarBuilder}; -use crate::mixed_precision::training_dtype; use ml_core::MLError; /// Noisy linear layer with factorized Gaussian noise (Rainbow DQN standard) @@ -66,7 +65,7 @@ impl NoisyLinear { let device = vb.device().clone(); // All init in F32, then cast — Candle's Init::Uniform CUDA kernel lacks BF16 PTX. - let dtype = training_dtype(&device); + let dtype = candle_core::DType::BF16; // Initialize μ_w ~ U(-1/√in, 1/√in) (Rainbow DQN standard) let mu_range = 1.0 / (in_features as f64).sqrt(); @@ -208,7 +207,7 @@ impl NoisyLinear { // Sample from N(0, 1), then cast to training dtype (BF16 on Ampere+ CUDA) let noise = Tensor::randn(0_f32, 1.0, size, device) .map_err(|e| MLError::ModelError(format!("Failed to sample noise: {}", e)))?; - let noise = noise.to_dtype(training_dtype(device)) + let noise = noise.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast noise to training dtype: {}", e)))?; // Apply f(x) = sign(x) × √|x| @@ -292,7 +291,7 @@ impl NoisyLinear { /// Disable noise for evaluation (use mean parameters only) pub fn disable_noise(&mut self) -> Result<(), MLError> { // Set epsilon buffers to zero in training dtype (effectively uses μ only) - let dtype = training_dtype(&self.device); + let dtype = candle_core::DType::BF16; self.weight_epsilon = Tensor::zeros((self.out_features, self.in_features), dtype, &self.device) .map_err(|e| MLError::ModelError(format!("Failed to zero weight_epsilon: {}", e)))?; self.bias_epsilon = Tensor::zeros(self.out_features, dtype, &self.device) @@ -332,13 +331,12 @@ impl Default for NoisyNetworkConfig { mod tests { use super::*; use candle_nn::{VarBuilder, VarMap}; - use crate::mixed_precision::training_dtype; #[test] fn test_noisy_linear_creation() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let _layer = NoisyLinear::new(64, 32, vb, 0.5)?; Ok(()) @@ -348,7 +346,7 @@ mod tests { fn test_noisy_linear_forward() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?; layer.reset_noise()?; // Resample noise before forward @@ -356,7 +354,7 @@ mod tests { // Create dummy input let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) .map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))? - .to_dtype(training_dtype(&device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?; // Forward pass @@ -372,12 +370,12 @@ mod tests { fn test_noise_reset() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?; let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) .map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))? - .to_dtype(training_dtype(&device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?; // First forward pass @@ -421,12 +419,12 @@ mod tests { fn test_disable_noise() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?; let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) .map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))? - .to_dtype(training_dtype(&device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?; // Reset noise for first pass @@ -465,7 +463,7 @@ mod tests { fn test_factorized_noise_dimensions() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let mut layer = NoisyLinear::new(128, 64, vb, 0.5)?; layer.reset_noise()?; @@ -481,12 +479,12 @@ mod tests { fn test_reset_noise_with_sigma() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?; let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) .map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))? - .to_dtype(training_dtype(&device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?; // Test with high sigma (0.6) @@ -525,7 +523,7 @@ mod tests { fn test_sigma_scaling_effect() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?; @@ -533,7 +531,7 @@ mod tests { layer.disable_noise()?; let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) .map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))? - .to_dtype(training_dtype(&device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?; let output_no_noise = layer.forward(&input)?; diff --git a/crates/ml-dqn/src/quantile_regression.rs b/crates/ml-dqn/src/quantile_regression.rs index ebaee8651..21a8c20ab 100644 --- a/crates/ml-dqn/src/quantile_regression.rs +++ b/crates/ml-dqn/src/quantile_regression.rs @@ -20,7 +20,6 @@ use candle_nn::{Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; use std::f32::consts::PI; -use crate::mixed_precision::training_dtype; use ml_core::MLError; /// Configuration for Quantile Regression DQN @@ -88,7 +87,7 @@ impl QuantileNetwork { vars: VarMap, device: &Device, ) -> Result { - let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, device); // Quantile embedding layer let quantile_embedding = candle_nn::linear( @@ -144,9 +143,9 @@ impl QuantileNetwork { /// # Returns /// Quantile values [batch, `num_actions`, `num_quantiles`] pub fn forward(&self, state_embed: &Tensor, taus: &Tensor) -> Result { - let state_embed = crate::mixed_precision::ensure_training_dtype(state_embed) + let state_embed = state_embed.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; - let taus = crate::mixed_precision::ensure_training_dtype(taus) + let taus = taus.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let batch_size = state_embed.dim(0)?; let num_quantiles = taus.dim(1)?; diff --git a/crates/ml-dqn/src/rainbow_agent.rs b/crates/ml-dqn/src/rainbow_agent.rs index 38bb1e922..43ff7a818 100644 --- a/crates/ml-dqn/src/rainbow_agent.rs +++ b/crates/ml-dqn/src/rainbow_agent.rs @@ -9,7 +9,6 @@ use std::sync::{Arc, Mutex, RwLock}; use ml_core::optimizers::Adam; use candle_core::{Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; -use crate::mixed_precision::training_dtype; use candle_optimisers::adam::ParamsAdam; use tracing::{debug, info}; @@ -57,10 +56,10 @@ impl RainbowAgent { let target_varmap = Arc::new(VarMap::new()); // Create networks - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let online_network = RainbowNetwork::new(&vs, config.network_config.clone())?; - let target_vs = VarBuilder::from_varmap(&target_varmap, training_dtype(&device), &device); + let target_vs = VarBuilder::from_varmap(&target_varmap, candle_core::DType::BF16, &device); let target_network = RainbowNetwork::new(&target_vs, config.network_config.clone())?; // Create optimizer diff --git a/crates/ml-dqn/src/rainbow_network.rs b/crates/ml-dqn/src/rainbow_network.rs index 6f320c50e..d69823957 100644 --- a/crates/ml-dqn/src/rainbow_network.rs +++ b/crates/ml-dqn/src/rainbow_network.rs @@ -453,13 +453,12 @@ mod tests { use anyhow::Result; use candle_core::Device; use candle_nn::{VarBuilder, VarMap}; - use crate::mixed_precision::training_dtype; #[test] fn test_rainbow_network_creation() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = RainbowNetworkConfig::default(); let _network = RainbowNetwork::new(&vs, config) @@ -480,7 +479,7 @@ mod tests { fn test_rainbow_activation_types() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let mut config = RainbowNetworkConfig::default(); config.activation = ActivationType::ReLU; diff --git a/crates/ml-dqn/src/replay_buffer_type.rs b/crates/ml-dqn/src/replay_buffer_type.rs index 08cc72bf4..aa47fae64 100644 --- a/crates/ml-dqn/src/replay_buffer_type.rs +++ b/crates/ml-dqn/src/replay_buffer_type.rs @@ -209,38 +209,11 @@ impl ReplayBufferType { })))) } - /// Try GPU-resident PER, fall back to CPU PER on allocation failure. - /// - /// On L40S/H100 this always succeeds (47 MB for 100K buffer). - /// On smaller GPUs or when VRAM is fragmented, gracefully degrades. - #[cfg(feature = "cuda")] - pub fn try_gpu_prioritized_with_fallback( - capacity: usize, - state_dim: usize, - alpha: f64, - beta: f64, - beta_max: f64, - beta_annealing_steps: usize, - max_memory_bytes: usize, - device: &candle_core::Device, - ) -> Result { - // Try GPU first, then CPU PER fallback on allocation failure. - match Self::try_gpu_with_halving(capacity, state_dim, alpha, beta, beta_max, beta_annealing_steps, max_memory_bytes, device) { - Ok(buf) => Ok(buf), - Err(e) => { - tracing::warn!( - "GPU PER allocation failed after retries ({}), falling back to CPU PER", - e - ); - Self::cpu_per_fallback(capacity, state_dim, alpha, beta, beta_max, beta_annealing_steps, max_memory_bytes, &e) - } - } - } - /// Attempt GPU PER allocation with adaptive capacity halving on OOM. - /// Returns the buffer on success, or the last error on exhaustion. + /// Returns the buffer on success, or a hard error on exhaustion. + /// No CPU PER fallback — GPU PER is mandatory on CUDA. #[cfg(feature = "cuda")] - fn try_gpu_with_halving( + pub fn try_gpu_with_halving( capacity: usize, state_dim: usize, alpha: f64, @@ -270,7 +243,7 @@ impl ReplayBufferType { return Ok(buf); } Err(e) => { - tracing::debug!( + tracing::warn!( "GPU PER at capacity {} failed ({}), retrying at {}", try_cap, e, try_cap / 2 ); @@ -285,44 +258,6 @@ impl ReplayBufferType { Err(MLError::ModelError("GPU PER capacity below floor".into())) } - /// CPU PER fallback after GPU allocation failure. - /// Pre-flight checks memory estimate against limit, then constructs CPU PER buffer. - #[cfg(feature = "cuda")] - fn cpu_per_fallback( - capacity: usize, - state_dim: usize, - alpha: f64, - beta: f64, - beta_max: f64, - beta_annealing_steps: usize, - max_memory_bytes: usize, - gpu_err: &MLError, - ) -> Result { - // SegmentTree: 2 * next_power_of_two(capacity) f32 entries - // PrioritizedReplayBuffer: capacity Option slots - let tree_elems = capacity - .checked_next_power_of_two() - .and_then(|p| p.checked_mul(2)) - .unwrap_or(usize::MAX); - let tree_bytes = tree_elems.saturating_mul(std::mem::size_of::()); - let exp_bytes = capacity.saturating_mul( - state_dim.saturating_mul(8).saturating_add(64), - ); - let estimated_bytes = tree_bytes.saturating_add(exp_bytes); - - if estimated_bytes > max_memory_bytes { - return Err(MLError::ModelError(format!( - "GPU PER failed ({}) and CPU PER fallback would need ~{} MB (limit {} MB). \ - Reduce capacity or state_dim.", - gpu_err, - estimated_bytes / (1024 * 1024), - max_memory_bytes / (1024 * 1024), - ))); - } - - Self::new_prioritized(capacity, alpha, beta, beta_max, beta_annealing_steps) - } - /// Sample a batch from the buffer pub fn sample(&self, batch_size: usize) -> Result { match self { diff --git a/crates/ml-dqn/src/residual.rs b/crates/ml-dqn/src/residual.rs index 9fdb12f90..88e78e517 100644 --- a/crates/ml-dqn/src/residual.rs +++ b/crates/ml-dqn/src/residual.rs @@ -108,7 +108,7 @@ impl ResidualBlock { /// Output tensor with same shape as input pub fn forward(&self, x: &Tensor, train: bool) -> Result { // Cast input to training dtype (BF16 on CUDA, F32 on CPU) - let x = crate::mixed_precision::ensure_training_dtype(x) + let x = x.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("dtype cast failed: {}", e)))?; // Save input for skip connection @@ -170,7 +170,6 @@ mod tests { use super::*; use candle_core::Device; use candle_nn::VarMap; - use crate::mixed_precision::training_dtype; #[test] fn test_residual_config_default() { @@ -184,7 +183,7 @@ mod tests { fn test_residual_block_creation() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let config = ResidualConfig { hidden_dim: 64, @@ -202,7 +201,7 @@ mod tests { fn test_residual_block_forward_train() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let config = ResidualConfig { hidden_dim: 32, @@ -214,7 +213,7 @@ mod tests { // Create input tensor (batch_size=2, hidden_dim=32) let input = Tensor::randn(0.0_f32, 1.0, (2, 32), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass in training mode let output = block.forward(&input, true)?; @@ -229,7 +228,7 @@ mod tests { fn test_residual_block_forward_eval() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let config = ResidualConfig { hidden_dim: 32, @@ -241,7 +240,7 @@ mod tests { // Create input tensor (batch_size=2, hidden_dim=32) let input = Tensor::randn(0.0_f32, 1.0, (2, 32), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass in eval mode (no dropout) let output = block.forward(&input, false)?; @@ -257,7 +256,7 @@ mod tests { // Test that skip connection preserves gradient flow let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let config = ResidualConfig { hidden_dim: 16, @@ -268,7 +267,7 @@ mod tests { let block = ResidualBlock::new(&var_builder, &config, "test_block")?; // Create simple input - let input = Tensor::ones((1, 16), training_dtype(&device), &device)?; + let input = Tensor::ones((1, 16), candle_core::DType::BF16, &device)?; // Forward pass let output = block.forward(&input, false)?; @@ -284,7 +283,7 @@ mod tests { fn test_residual_batch_processing() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let config = ResidualConfig { hidden_dim: 64, @@ -297,7 +296,7 @@ mod tests { // Test different batch sizes for batch_size in [1, 4, 8, 16] { let input = Tensor::randn(0.0_f32, 1.0, (batch_size, 64), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; let output = block.forward(&input, true)?; assert_eq!(output.dims(), &[batch_size, 64]); } @@ -310,7 +309,7 @@ mod tests { // Test that gradients can flow through skip connection let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let config = ResidualConfig { hidden_dim: 8, @@ -322,7 +321,7 @@ mod tests { // Create input with requires_grad let input = Tensor::randn(0.0_f32, 1.0, (2, 8), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass let output = block.forward(&input, false)?; @@ -340,7 +339,7 @@ mod tests { fn test_residual_different_dimensions() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); // Test different hidden dimensions for hidden_dim in [16, 32, 64, 128, 256] { @@ -352,7 +351,7 @@ mod tests { let block = ResidualBlock::new(&var_builder.pp(format!("block_{}", hidden_dim)), &config, "test")?; let input = Tensor::randn(0.0_f32, 1.0, (2, hidden_dim), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; let output = block.forward(&input, true)?; assert_eq!(output.dims(), &[2, hidden_dim]); } @@ -364,7 +363,7 @@ mod tests { fn test_residual_numerical_stability() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let config = ResidualConfig { hidden_dim: 32, @@ -376,7 +375,7 @@ mod tests { // Test with extreme values let input = Tensor::from_vec(vec![100.0_f32; 32], (1, 32), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; let output = block.forward(&input, false)?; // Check for NaN/Inf diff --git a/crates/ml-dqn/src/rmsnorm.rs b/crates/ml-dqn/src/rmsnorm.rs index 2de6d136f..525cca4d2 100644 --- a/crates/ml-dqn/src/rmsnorm.rs +++ b/crates/ml-dqn/src/rmsnorm.rs @@ -244,14 +244,13 @@ mod tests { use candle_core::Device; use candle_nn::VarMap; use std::time::Instant; - use crate::mixed_precision::training_dtype; use tracing::info; #[test] fn test_rmsnorm_creation() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let dim = 128; let rmsnorm = RMSNorm::new_default(vs.pp("rmsnorm"), dim)?; @@ -266,7 +265,7 @@ mod tests { fn test_layernorm_creation() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let dim = 128; let layernorm = LayerNorm::new_default(vs.pp("layernorm"), dim)?; @@ -281,7 +280,7 @@ mod tests { fn test_rmsnorm_forward() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let batch_size = 4; let dim = 128; @@ -292,7 +291,7 @@ mod tests { .map(|i| (i as f32 * 0.01).sin()) .collect(); let input = Tensor::from_vec(input_data, (batch_size, dim), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass let output = rmsnorm.forward(&input)?; @@ -323,7 +322,7 @@ mod tests { fn test_layernorm_forward() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let batch_size = 4; let dim = 128; @@ -334,7 +333,7 @@ mod tests { .map(|i| (i as f32 * 0.01).sin()) .collect(); let input = Tensor::from_vec(input_data, (batch_size, dim), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass let output = layernorm.forward(&input)?; @@ -369,12 +368,12 @@ mod tests { // Setup RMSNorm let rmsnorm_varmap = VarMap::new(); - let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, training_dtype(&device), &device); + let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, candle_core::DType::BF16, &device); let rmsnorm = RMSNorm::new_default(rmsnorm_vs.pp("rmsnorm"), dim)?; // Setup LayerNorm let layernorm_varmap = VarMap::new(); - let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, training_dtype(&device), &device); + let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, candle_core::DType::BF16, &device); let layernorm = LayerNorm::new_default(layernorm_vs.pp("layernorm"), dim)?; // Create random input @@ -382,7 +381,7 @@ mod tests { .map(|i| (i as f32 * 0.01).sin()) .collect(); let input = Tensor::from_vec(input_data, (batch_size, dim), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Benchmark RMSNorm let rmsnorm_start = Instant::now(); @@ -425,11 +424,11 @@ mod tests { // Setup both norms let rmsnorm_varmap = VarMap::new(); - let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, training_dtype(&device), &device); + let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, candle_core::DType::BF16, &device); let rmsnorm = RMSNorm::new_default(rmsnorm_vs.pp("rmsnorm"), dim)?; let layernorm_varmap = VarMap::new(); - let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, training_dtype(&device), &device); + let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, candle_core::DType::BF16, &device); let layernorm = LayerNorm::new_default(layernorm_vs.pp("layernorm"), dim)?; // Create random input @@ -437,7 +436,7 @@ mod tests { .map(|i| (i as f32 * 0.01).sin()) .collect(); let input = Tensor::from_vec(input_data, (batch_size, dim), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass through both let rmsnorm_output = rmsnorm.forward(&input)?; @@ -480,7 +479,7 @@ mod tests { fn test_rmsnorm_3d_input() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let batch_size = 4; let seq_len = 16; @@ -492,7 +491,7 @@ mod tests { .map(|i| (i as f32 * 0.01).sin()) .collect(); let input = Tensor::from_vec(input_data, (batch_size, seq_len, dim), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass let output = rmsnorm.forward(&input)?; @@ -507,7 +506,7 @@ mod tests { fn test_layernorm_3d_input() -> anyhow::Result<()> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let batch_size = 4; let seq_len = 16; @@ -519,7 +518,7 @@ mod tests { .map(|i| (i as f32 * 0.01).sin()) .collect(); let input = Tensor::from_vec(input_data, (batch_size, seq_len, dim), &device)? - .to_dtype(training_dtype(&device))?; + .to_dtype(candle_core::DType::BF16)?; // Forward pass let output = layernorm.forward(&input)?; diff --git a/crates/ml-dqn/tests/gpu_smoketest.rs b/crates/ml-dqn/tests/gpu_smoketest.rs index 22a798442..e6a60ecb8 100644 --- a/crates/ml-dqn/tests/gpu_smoketest.rs +++ b/crates/ml-dqn/tests/gpu_smoketest.rs @@ -54,7 +54,7 @@ fn smoketest_config() -> DQNConfig { n_steps: 1, initial_capital: 100_000.0, use_per: true, // GPU PER mandatory on CUDA - use_gpu_replay_buffer: true, // GPU PER active + per_alpha: 0.6, per_beta_start: 0.4, per_beta_max: 1.0, @@ -99,7 +99,6 @@ fn smoketest_config() -> DQNConfig { minimum_profit_factor: 1.5, weight_decay: 0.0, dropout_rate: 0.0, - mixed_precision: None, } } diff --git a/crates/ml-explainability/src/integrated_gradients.rs b/crates/ml-explainability/src/integrated_gradients.rs index 7ae053353..4c17396d0 100644 --- a/crates/ml-explainability/src/integrated_gradients.rs +++ b/crates/ml-explainability/src/integrated_gradients.rs @@ -198,7 +198,7 @@ mod tests { fn test_integrated_gradients_basic() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, ml_core::mixed_precision::training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let model = TwoLayerNet::new(vs); // Linear layers require 2D input: [batch, features] @@ -246,7 +246,7 @@ mod tests { fn test_ig_completeness_axiom() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, ml_core::mixed_precision::training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let model = LinearNet::new(vs); let input = Tensor::new(&[[1.0_f32, -0.5, 0.3, 2.0]], &device).unwrap(); @@ -301,7 +301,7 @@ mod tests { fn test_ig_dimension_mismatch() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, ml_core::mixed_precision::training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let model = TwoLayerNet::new(vs); let input = Tensor::new(&[[1.0_f32, -0.5, 0.3, 2.0]], &device).unwrap(); diff --git a/crates/ml-ppo/src/continuous_demo.rs b/crates/ml-ppo/src/continuous_demo.rs index b4d8666d0..4ff7de8e9 100644 --- a/crates/ml-ppo/src/continuous_demo.rs +++ b/crates/ml-ppo/src/continuous_demo.rs @@ -52,7 +52,7 @@ pub fn demo_continuous_position_sizing() -> Result<(), MLError> { for (scenario_name, state_vec) in market_scenarios { let state_tensor = - Tensor::from_vec(state_vec, (1, 8), &device)?.to_dtype(ml_core::mixed_precision::training_dtype(&device))?; + Tensor::from_vec(state_vec, (1, 8), &device)?.to_dtype(candle_core::DType::BF16)?; // Sample multiple actions to show distribution let mut position_sizes = Vec::new(); @@ -90,7 +90,7 @@ pub fn demo_continuous_position_sizing() -> Result<(), MLError> { // Show entropy (exploration level) let test_state = - Tensor::from_vec(vec![0.5; 8], (1, 8), &device)?.to_dtype(ml_core::mixed_precision::training_dtype(&device))?; + Tensor::from_vec(vec![0.5; 8], (1, 8), &device)?.to_dtype(candle_core::DType::BF16)?; let entropy = policy.entropy(&test_state)?; let entropy_value = entropy.flatten_all()?.to_dtype(candle_core::DType::F32)?.to_scalar::()?; @@ -176,7 +176,7 @@ pub fn trading_integration_example() -> Result<(), MLError> { ]; let state_tensor = - Tensor::from_vec(trading_state, (1, 16), &device)?.to_dtype(ml_core::mixed_precision::training_dtype(&device))?; + Tensor::from_vec(trading_state, (1, 16), &device)?.to_dtype(candle_core::DType::BF16)?; // Get position sizing recommendation let (action_value, log_prob) = policy.sample_action(&state_tensor)?; diff --git a/crates/ml-ppo/src/continuous_policy.rs b/crates/ml-ppo/src/continuous_policy.rs index e0816b945..94c9447e2 100644 --- a/crates/ml-ppo/src/continuous_policy.rs +++ b/crates/ml-ppo/src/continuous_policy.rs @@ -21,7 +21,6 @@ use serde::{Deserialize, Serialize}; use statrs::distribution::{ContinuousCDF, Normal}; use tracing::{debug, warn}; -use ml_core::mixed_precision::training_dtype; use ml_core::xavier_init::linear_xavier; use ml_core::MLError; @@ -81,7 +80,7 @@ impl ContinuousPolicyNetwork { /// Create new continuous policy network pub fn new(config: ContinuousPolicyConfig, device: Device) -> Result { let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let mut feature_layers = Vec::new(); let mut current_dim = config.state_dim; @@ -113,10 +112,9 @@ impl ContinuousPolicyNetwork { })?; (Some(log_std_head), None) } else { - let dtype = training_dtype(&device); let fixed_log_std = Tensor::full(config.init_log_std, (1, 1), &device) - .and_then(|t| t.to_dtype(dtype)) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| { MLError::ModelError(format!("Failed to create fixed log std: {}", e)) })?; @@ -180,10 +178,9 @@ impl ContinuousPolicyNetwork { })?; (Some(log_std_head), None) } else { - let dtype = training_dtype(&device); let fixed_log_std = Tensor::full(config.init_log_std, (1, 1), &device) - .and_then(|t| t.to_dtype(dtype)) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| { MLError::ModelError(format!("Failed to create fixed log std: {}", e)) })?; @@ -208,7 +205,7 @@ impl ContinuousPolicyNetwork { /// Forward pass returning mean and log standard deviation pub fn forward(&self, input: &Tensor) -> Result<(Tensor, Tensor), MLError> { let mut x = input - .to_dtype(ml_core::mixed_precision::training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Input dtype cast failed: {}", e)))?; // Pass through shared feature layers @@ -450,10 +447,9 @@ impl ContinuousPolicyNetwork { let max_log_std = self.config.max_log_std.min(2.0); let clamped_log_std = log_std.clamp(min_log_std, max_log_std); - let dtype = training_dtype(&self.device); self.fixed_log_std = Some( Tensor::full(clamped_log_std, (1, 1), &self.device) - .and_then(|t| t.to_dtype(dtype)) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| MLError::ModelError(format!("Failed to set log std: {}", e)))?, ); diff --git a/crates/ml-ppo/src/continuous_ppo.rs b/crates/ml-ppo/src/continuous_ppo.rs index 04a67d04c..b990426a4 100644 --- a/crates/ml-ppo/src/continuous_ppo.rs +++ b/crates/ml-ppo/src/continuous_ppo.rs @@ -4,7 +4,6 @@ //! action spaces, using Gaussian policies for position sizing. use candle_core::{DType, Device, Tensor}; -use ml_core::mixed_precision::training_dtype; use candle_nn::Optimizer; // Required for Adam::new and backward_step methods use candle_optimisers::adam::Adam; use candle_optimisers::adam::ParamsAdam; @@ -205,7 +204,7 @@ impl ContinuousTrajectoryBatch { state_dim: usize, ) -> Result { let batch_size = self.states.len(); - let dtype = training_dtype(device); + let dtype = candle_core::DType::BF16; let state_flat: Vec = self.states.iter().flatten().cloned().collect(); let states = Tensor::from_vec(state_flat, (batch_size, state_dim), device) @@ -289,7 +288,7 @@ impl ContinuousMiniBatch { state_dim: usize, ) -> Result { let batch_size = self.states.len(); - let dtype = training_dtype(device); + let dtype = candle_core::DType::BF16; let state_flat: Vec = self.states.iter().flatten().cloned().collect(); let states = Tensor::from_vec(state_flat, (batch_size, state_dim), device) @@ -393,7 +392,7 @@ impl ContinuousPPO { (1, self.config.state_dim), self.actor.device(), ) - .and_then(|t| t.to_dtype(training_dtype(self.actor.device()))) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))?; // Get action from policy @@ -427,7 +426,7 @@ impl ContinuousPPO { (1, self.config.state_dim), self.actor.device(), ) - .and_then(|t| t.to_dtype(training_dtype(self.actor.device()))) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))?; // Get action and log prob from policy @@ -584,7 +583,7 @@ impl ContinuousPPO { let ratio = clipped_log_ratio.exp()?; // Clipped surrogate objective - let dtype = training_dtype(self.actor.device()); + let dtype = candle_core::DType::BF16; let clip_epsilon_tensor = Tensor::from_vec( vec![self.config.clip_epsilon; batch.advantages.dims()[0]], batch.advantages.dims(), diff --git a/crates/ml-ppo/src/flow_policy/mod.rs b/crates/ml-ppo/src/flow_policy/mod.rs index 483067b1d..c917a3ed6 100644 --- a/crates/ml-ppo/src/flow_policy/mod.rs +++ b/crates/ml-ppo/src/flow_policy/mod.rs @@ -10,7 +10,6 @@ use rand::thread_rng; use rand_distr::{Distribution, Normal}; use serde::{Deserialize, Serialize}; -use ml_core::mixed_precision::training_dtype; use ml_core::xavier_init::linear_xavier; use ml_core::MLError; @@ -123,7 +122,7 @@ impl FlowPolicy { /// A new `FlowPolicy` instance or an error if initialization fails. pub fn new(config: FlowPolicyConfig, device: &Device) -> Result { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, device); // Context encoder: state_dim → context_dim let context_enc = linear_xavier( @@ -330,7 +329,7 @@ impl FlowPolicy { let ctx = self.encode_context(states)?; // Cast actions to training dtype to match flow layer weights - let actions = actions.to_dtype(training_dtype(&self.device)) + let actions = actions.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::TensorOperationError(format!("Actions dtype cast failed: {}", e)))?; // Unsquash: a → y = atanh(a) = 0.5 * ln((1+a)/(1-a)) @@ -414,7 +413,7 @@ impl FlowPolicy { /// # Returns /// Context tensor [`batch_size`, `context_dim`]. fn encode_context(&self, state: &Tensor) -> Result { - let state = state.to_dtype(training_dtype(&self.device)) + let state = state.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::TensorOperationError(format!("State dtype cast failed: {}", e)))?; let h = self.context_enc.forward(&state).map_err(|e| { MLError::TensorOperationError(format!("Context encoder forward failed: {}", e)) @@ -434,7 +433,7 @@ impl FlowPolicy { fn flow_forward(&self, z: &Tensor, ctx: &Tensor) -> Result<(Tensor, Tensor), MLError> { let mut x = z.clone(); let batch_size = z.dims()[0]; - let mut log_det_acc = Tensor::zeros(batch_size, training_dtype(&self.device), &self.device) + let mut log_det_acc = Tensor::zeros(batch_size, candle_core::DType::BF16, &self.device) .map_err(|e| MLError::TensorOperationError(format!("Log det init failed: {}", e)))?; for layer in &self.layers { @@ -459,7 +458,7 @@ impl FlowPolicy { fn flow_inverse(&self, y: &Tensor, ctx: &Tensor) -> Result<(Tensor, Tensor), MLError> { let mut x = y.clone(); let batch_size = y.dims()[0]; - let mut log_det_acc = Tensor::zeros(batch_size, training_dtype(&self.device), &self.device) + let mut log_det_acc = Tensor::zeros(batch_size, candle_core::DType::BF16, &self.device) .map_err(|e| MLError::TensorOperationError(format!("Log det init failed: {}", e)))?; // Reverse order of layers for inverse @@ -492,7 +491,7 @@ impl FlowPolicy { .collect(); Tensor::from_vec(samples, (batch_size, self.config.action_dim), &self.device) - .and_then(|t| t.to_dtype(training_dtype(&self.device))) + .and_then(|t| t.to_dtype(candle_core::DType::BF16)) .map_err(|e| MLError::TensorOperationError(format!("Noise tensor creation failed: {}", e))) } @@ -532,7 +531,7 @@ impl FlowPolicy { /// Normalizing flows don't have a fixed `log_std` parameter like Gaussian policies. /// Returns zeros for API compatibility with existing PPO code. pub fn get_current_log_std(&self) -> Result { - Tensor::zeros(self.config.action_dim, training_dtype(&self.device), &self.device) + Tensor::zeros(self.config.action_dim, candle_core::DType::BF16, &self.device) .map_err(|e| MLError::TensorOperationError(format!("Log std creation failed: {}", e))) } diff --git a/crates/ml-ppo/src/hidden_state_manager.rs b/crates/ml-ppo/src/hidden_state_manager.rs index 9845a4f3e..574918c30 100644 --- a/crates/ml-ppo/src/hidden_state_manager.rs +++ b/crates/ml-ppo/src/hidden_state_manager.rs @@ -8,7 +8,6 @@ use candle_core::{Device, Tensor}; use candle_core::DType; use std::fmt; use ml_core::MLError; -use ml_core::mixed_precision::training_dtype; /// Manages LSTM hidden and cell states for policy and value networks pub struct HiddenStateManager { @@ -43,7 +42,7 @@ impl HiddenStateManager { device: &Device, ) -> Result { let shape = &[num_layers, batch_size, hidden_dim]; - let zeros = Tensor::zeros(shape, training_dtype(device), device) + let zeros = Tensor::zeros(shape, candle_core::DType::BF16, device) .map_err(|e| MLError::TensorOperationError(format!("Failed to create zero tensor: {}", e)))?; Ok(Self { @@ -140,9 +139,8 @@ impl HiddenStateManager { // Convert done_mask to float and expand to match state dimensions // done_mask: [batch_size] -> [1, batch_size, 1] - let dtype = training_dtype(&self.device); let done_float = done_mask - .to_dtype(dtype) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::TensorOperationError(format!("Failed to convert done mask to float: {}", e)))?; let done_expanded = done_float @@ -157,7 +155,7 @@ impl HiddenStateManager { .map_err(|e| MLError::TensorOperationError(format!("Failed to broadcast done mask: {}", e)))?; // Create keep_mask = 1 - done_mask (keep states where episode continues) - let ones = Tensor::ones(&[self.num_layers, self.batch_size, self.hidden_dim], dtype, &self.device) + let ones = Tensor::ones(&[self.num_layers, self.batch_size, self.hidden_dim], candle_core::DType::BF16, &self.device) .map_err(|e| MLError::TensorOperationError(format!("Failed to create ones tensor: {}", e)))?; let keep_mask = ones @@ -187,7 +185,7 @@ impl HiddenStateManager { /// Reset all states to zeros pub fn reset_all(&mut self) -> Result<(), MLError> { let shape = &[self.num_layers, self.batch_size, self.hidden_dim]; - let zeros = Tensor::zeros(shape, training_dtype(&self.device), &self.device) + let zeros = Tensor::zeros(shape, candle_core::DType::BF16, &self.device) .map_err(|e| MLError::TensorOperationError(format!("Failed to create zero tensor: {}", e)))?; self.policy_hidden = zeros.clone(); @@ -242,9 +240,8 @@ mod tests { let device = cuda_device(); let mut manager = HiddenStateManager::new(1, 2, 3, &device)?; - let dtype = training_dtype(&device); - let new_h = Tensor::ones(&[1, 2, 3], dtype, &device)?; - let new_c = Tensor::ones(&[1, 2, 3], dtype, &device)?; + let new_h = Tensor::ones(&[1, 2, 3], candle_core::DType::BF16, &device)?; + let new_c = Tensor::ones(&[1, 2, 3], candle_core::DType::BF16, &device)?; manager.update_policy_state(new_h.clone(), new_c.clone())?; @@ -264,8 +261,7 @@ mod tests { let mut manager = HiddenStateManager::new(1, 2, 3, &device)?; // Set to non-zero - let dtype = training_dtype(&device); - let ones = Tensor::ones(&[1, 2, 3], dtype, &device)?; + let ones = Tensor::ones(&[1, 2, 3], candle_core::DType::BF16, &device)?; manager.update_policy_state(ones.clone(), ones.clone())?; manager.update_value_state(ones.clone(), ones)?; diff --git a/crates/ml-ppo/src/lstm_networks.rs b/crates/ml-ppo/src/lstm_networks.rs index bcae87df3..db28f62ea 100644 --- a/crates/ml-ppo/src/lstm_networks.rs +++ b/crates/ml-ppo/src/lstm_networks.rs @@ -9,7 +9,6 @@ use candle_core::{Device, Tensor}; use candle_nn::{linear, Linear, Module, VarBuilder, VarMap, LSTM, LSTMConfig}; use candle_nn::rnn::{RNN, LSTMState}; -use ml_core::mixed_precision::training_dtype; use ml_core::MLError; /// LSTM-augmented policy network for temporal action selection @@ -46,7 +45,7 @@ impl LSTMPolicyNetwork { device: Device, ) -> Result { let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); // Input projection layer (state_dim → hidden_dim) let input_layer = linear(input_dim, hidden_dim, var_builder.pp("input")) @@ -98,7 +97,7 @@ impl LSTMPolicyNetwork { h_t: &Tensor, c_t: &Tensor, ) -> Result<(Tensor, Tensor, Tensor), MLError> { - let state = ml_core::mixed_precision::ensure_training_dtype(state) + let state = state.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; // Project input: [batch, input_dim] → [batch, hidden_dim] let x = self @@ -252,7 +251,7 @@ impl LSTMValueNetwork { device: Device, ) -> Result { let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); // Input projection layer (state_dim → hidden_dim) let input_layer = linear(input_dim, hidden_dim, var_builder.pp("input")) @@ -303,7 +302,7 @@ impl LSTMValueNetwork { h_t: &Tensor, c_t: &Tensor, ) -> Result<(Tensor, Tensor, Tensor), MLError> { - let state = ml_core::mixed_precision::ensure_training_dtype(state) + let state = state.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; // Project input: [batch, input_dim] → [batch, hidden_dim] let x = self diff --git a/crates/ml-ppo/src/ppo.rs b/crates/ml-ppo/src/ppo.rs index d3ce16091..f6b67103b 100644 --- a/crates/ml-ppo/src/ppo.rs +++ b/crates/ml-ppo/src/ppo.rs @@ -28,7 +28,6 @@ use super::hidden_state_manager::HiddenStateManager; use super::lstm_networks::{LSTMPolicyNetwork, LSTMValueNetwork}; use super::trajectories::{TrajectoryBatch, TrajectoryTensors}; use ml_core::common::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig}; -use ml_core::mixed_precision::training_dtype; use ml_core::portfolio_tracker::PortfolioTracker; use ml_core::xavier_init::linear_xavier; use crate::reward_normalizer::RewardNormalizer; @@ -240,9 +239,6 @@ pub struct PPOConfig { /// When Some, uses clip(ratio, 1-clip_epsilon, `1+clip_epsilon_high`). /// Prevents entropy collapse during long training. None = symmetric (default). pub clip_epsilon_high: Option, - /// Mixed precision configuration for BF16/FP16 forward pass on supported GPUs. - /// None = FP32 only. Auto-configured based on GPU architecture at runtime. - pub mixed_precision: Option, /// Use symlog transform for value targets (`DreamerV3`). Default: true. /// Compresses large returns while preserving sign. pub use_symlog: bool, @@ -285,7 +281,6 @@ impl Default for PPOConfig { lstm_sequence_length: 32, accumulation_steps: 1, clip_epsilon_high: Some(0.28), // DAPO asymmetric clipping: [1-0.2, 1+0.28] = [0.8, 1.28] - mixed_precision: None, use_symlog: true, use_adaptive_entropy: true, use_percentile_scaling: true, @@ -299,7 +294,6 @@ pub struct PolicyNetwork { layers: Vec, device: Device, vars: VarMap, - mixed_precision: Option, } impl PolicyNetwork { @@ -311,7 +305,7 @@ impl PolicyNetwork { device: Device, ) -> Result { let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let mut layers = Vec::new(); let mut current_dim = input_dim; @@ -343,7 +337,6 @@ impl PolicyNetwork { layers, device, vars, - mixed_precision: None, }) } @@ -418,39 +411,13 @@ impl PolicyNetwork { layers, device, vars, - mixed_precision: None, }) } /// Forward pass returning action logits pub fn forward(&self, input: &Tensor) -> Result { - self.forward_mixed(input, &self.mixed_precision) - } - - /// Set mixed precision config for this network - pub const fn set_mixed_precision(&mut self, mp: Option) { - self.mixed_precision = mp; - } - - /// Forward pass with optional mixed precision (BF16/FP16). - /// Casts input to reduced precision for compute, casts output back to FP32. - pub fn forward_mixed( - &self, - input: &Tensor, - mixed_precision: &Option, - ) -> Result { - let input = ml_core::mixed_precision::ensure_training_dtype(input) + let mut x = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; - let (mut x, _use_amp) = match mixed_precision { - Some(mp) if mp.enabled => { - let target_dtype = mp.dtype.to_dtype(); - match input.to_dtype(target_dtype) { - Ok(converted) => (converted, true), - Err(_) => (input, false), - } - } - _ => (input, false), - }; for (i, layer) in self.layers.iter().enumerate() { x = layer.forward(&x).map_err(|e| { @@ -569,14 +536,13 @@ pub struct ValueNetwork { layers: Vec, device: Device, vars: VarMap, - mixed_precision: Option, } impl ValueNetwork { /// Create new value network pub fn new(input_dim: usize, hidden_dims: &[usize], device: Device) -> Result { let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); + let var_builder = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device); let mut layers = Vec::new(); let mut current_dim = input_dim; @@ -608,7 +574,6 @@ impl ValueNetwork { layers, device, vars, - mixed_precision: None, }) } @@ -680,38 +645,13 @@ impl ValueNetwork { layers, device, vars, - mixed_precision: None, }) } /// Forward pass returning state values pub fn forward(&self, input: &Tensor) -> Result { - self.forward_mixed(input, &self.mixed_precision) - } - - /// Set mixed precision config for this network - pub const fn set_mixed_precision(&mut self, mp: Option) { - self.mixed_precision = mp; - } - - /// Forward pass with optional mixed precision (BF16/FP16). - pub fn forward_mixed( - &self, - input: &Tensor, - mixed_precision: &Option, - ) -> Result { - let input = ml_core::mixed_precision::ensure_training_dtype(input) + let mut x = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; - let (mut x, _use_amp) = match mixed_precision { - Some(mp) if mp.enabled => { - let target_dtype = mp.dtype.to_dtype(); - match input.to_dtype(target_dtype) { - Ok(converted) => (converted, true), - Err(_) => (input, false), - } - } - _ => (input, false), - }; for (i, layer) in self.layers.iter().enumerate() { x = layer.forward(&x).map_err(|e| { @@ -851,25 +791,19 @@ impl PPO { ) } else { // MLP mode (default): Create standard feedforward networks - let mut mlp_actor = PolicyNetwork::new( + let mlp_actor = PolicyNetwork::new( config.state_dim, &config.policy_hidden_dims, config.num_actions, device.clone(), )?; - let mut mlp_critic = ValueNetwork::new( + let mlp_critic = ValueNetwork::new( config.state_dim, &config.value_hidden_dims, device, )?; - // Wire mixed precision (BF16/FP16) into forward passes if configured - if config.mixed_precision.is_some() { - mlp_actor.set_mixed_precision(config.mixed_precision.clone()); - mlp_critic.set_mixed_precision(config.mixed_precision.clone()); - } - ( ActorNetwork::MLP(mlp_actor), CriticNetwork::MLP(mlp_critic), @@ -1484,18 +1418,17 @@ impl PPO { // Upload sequence metadata tensors (old log probs, advantages, returns). // Cast to training dtype so arithmetic with BF16 network outputs // doesn't trigger dtype mismatch errors. - let train_dt = ml_core::mixed_precision::training_dtype(device); let seq_old_log_probs = Tensor::from_vec( sequence.log_probs.clone(), (seq_len,), device, - )?.to_dtype(train_dt)?; + )?.to_dtype(candle_core::DType::BF16)?; let seq_advantages = Tensor::from_vec( sequence.advantages.clone(), (seq_len,), device, - )?.to_dtype(train_dt)?; + )?.to_dtype(candle_core::DType::BF16)?; let seq_returns = Tensor::from_vec( sequence.returns.clone(), @@ -2055,7 +1988,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], training_dtype(&device), &device).map_err( + VarBuilder::from_mmaped_safetensors(&[actor_path], candle_core::DType::BF16, &device).map_err( |e| { MLError::ModelError(format!( "Failed to load actor checkpoint from {}: {}", @@ -2112,7 +2045,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], training_dtype(&device), &device).map_err( + VarBuilder::from_mmaped_safetensors(&[critic_path], candle_core::DType::BF16, &device).map_err( |e| { MLError::ModelError(format!( "Failed to load critic checkpoint from {}: {}", diff --git a/crates/ml-ppo/src/trajectories.rs b/crates/ml-ppo/src/trajectories.rs index c7b7085bc..e96077c85 100644 --- a/crates/ml-ppo/src/trajectories.rs +++ b/crates/ml-ppo/src/trajectories.rs @@ -7,7 +7,6 @@ use candle_core::Tensor; use serde::{Deserialize, Serialize}; use ml_core::action_space::FactoredAction; -use ml_core::mixed_precision::training_dtype; use ml_core::MLError; /// Single step trajectory data @@ -288,7 +287,7 @@ impl TrajectoryBatch { state_dim: usize, ) -> Result { let batch_size = self.total_steps(); - let dtype = training_dtype(device); + let dtype = candle_core::DType::BF16; // Use pre-flattened states if available (from from_trajectories); // fall back to flatten_states_slow for GPU-constructed batches where states_flat is empty. @@ -532,7 +531,7 @@ impl MiniBatch { state_dim: usize, ) -> Result { let batch_size = self.states.len(); - let dtype = training_dtype(device); + let dtype = candle_core::DType::BF16; // Flatten states via extend_from_slice (contiguous memcpy per state vector) let states_flat = self.flatten_states(state_dim); diff --git a/crates/ml-supervised/src/diffusion/denoiser.rs b/crates/ml-supervised/src/diffusion/denoiser.rs index 0c3abaab0..7b989bf25 100644 --- a/crates/ml-supervised/src/diffusion/denoiser.rs +++ b/crates/ml-supervised/src/diffusion/denoiser.rs @@ -64,7 +64,7 @@ impl TimeEmbedding { .map_err(|e| MLError::ModelError(e.to_string()))?; // Cast to training dtype before projection through BF16 weights - let emb = ml_core::mixed_precision::ensure_training_dtype(&emb) + let emb = emb.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; // Project to hidden_dim @@ -203,7 +203,7 @@ impl Denoiser { /// /// Input x: (batch, `data_dim`), t: (batch,) → Output: (batch, `data_dim`) pub fn forward(&self, x: &Tensor, t: &Tensor) -> Result { - let x = ml_core::mixed_precision::ensure_training_dtype(x) + let x = x.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let map_err = |e: candle_core::Error| MLError::ModelError(e.to_string()); @@ -231,14 +231,14 @@ impl Denoiser { mod tests { use super::*; use candle_nn::VarMap; - use ml_core::mixed_precision::training_dtype; + #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_time_embedding_shape() { let dev = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let te = TimeEmbedding::new(32, 64, vb).unwrap(); let t = Tensor::new(&[0_u32, 100, 500, 999], &dev).unwrap(); let emb = te.forward(&t, &dev).unwrap(); @@ -250,7 +250,7 @@ mod tests { fn test_time_embedding_different_timesteps_differ() { let dev = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let te = TimeEmbedding::new(32, 64, vb).unwrap(); let t1 = Tensor::new(&[0_u32], &dev).unwrap(); let t2 = Tensor::new(&[500_u32], &dev).unwrap(); @@ -265,7 +265,7 @@ mod tests { fn test_denoiser_output_shape() { let dev = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let denoiser = Denoiser::new(64, 128, 3, 32, vb, &dev).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[4, 64], &dev).unwrap(); let t = Tensor::new(&[100_u32, 200, 300, 400], &dev).unwrap(); @@ -278,7 +278,7 @@ mod tests { fn test_denoiser_produces_gradients() { let dev = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let denoiser = Denoiser::new(64, 128, 2, 32, vb, &dev).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[4, 64], &dev).unwrap(); let t = Tensor::new(&[50_u32, 100, 200, 300], &dev).unwrap(); @@ -294,7 +294,7 @@ mod tests { fn test_denoiser_single_layer() { let dev = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let denoiser = Denoiser::new(32, 64, 1, 16, vb, &dev).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[2, 32], &dev).unwrap(); let t = Tensor::new(&[0_u32, 999], &dev).unwrap(); diff --git a/crates/ml-supervised/src/diffusion/sampler.rs b/crates/ml-supervised/src/diffusion/sampler.rs index e14de9204..cd56c3efe 100644 --- a/crates/ml-supervised/src/diffusion/sampler.rs +++ b/crates/ml-supervised/src/diffusion/sampler.rs @@ -145,7 +145,7 @@ mod tests { use super::*; use super::super::config::{DiffusionConfig, NoiseSchedule}; use candle_nn::{VarBuilder, VarMap}; - use ml_core::mixed_precision::training_dtype; + fn make_test_components() -> (Denoiser, NoiseScheduler, DDIMSampler) { let dev = Device::new_cuda(0).expect("CUDA required"); @@ -161,7 +161,7 @@ mod tests { }; let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let denoiser = Denoiser::new( config.data_dim(), config.hidden_dim, @@ -235,7 +235,7 @@ mod tests { fn test_ddim_single_step() { let dev = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let denoiser = Denoiser::new(8, 16, 1, 8, vb, &dev).unwrap(); let scheduler = NoiseScheduler::new(10, &NoiseSchedule::Linear, &dev).unwrap(); let sampler = DDIMSampler::new(1, 10, 0.0); diff --git a/crates/ml-supervised/src/kan/layer.rs b/crates/ml-supervised/src/kan/layer.rs index bcc71ab6f..ffa5146fc 100644 --- a/crates/ml-supervised/src/kan/layer.rs +++ b/crates/ml-supervised/src/kan/layer.rs @@ -143,14 +143,13 @@ mod tests { use super::*; use candle_core::Device; use candle_nn::VarMap; - use ml_core::mixed_precision::training_dtype; #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_kan_layer_output_shape() { let device = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &device); let layer = KANLayer::new(4, 8, 5, 4, vb.pp("layer0")).unwrap(); let input = Tensor::randn(0.0_f32, 0.5, &[2, 4], &device).unwrap(); @@ -164,7 +163,7 @@ mod tests { fn test_kan_layer_learnable_params() { let device = Device::new_cuda(0).expect("CUDA required"); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &device); let _layer = KANLayer::new(4, 8, 5, 4, vb.pp("layer0")).unwrap(); let all_vars = var_map.all_vars(); diff --git a/crates/ml-supervised/src/kan/network.rs b/crates/ml-supervised/src/kan/network.rs index ac6314937..311b7f206 100644 --- a/crates/ml-supervised/src/kan/network.rs +++ b/crates/ml-supervised/src/kan/network.rs @@ -54,7 +54,7 @@ impl KANNetwork { /// Input shape: `(batch, layer_widths[0])` /// Output shape: `(batch, layer_widths[last])` pub fn forward(&self, input: &Tensor) -> Result { - let input = ml_core::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let mut x = input; for layer in &self.layers { @@ -70,7 +70,6 @@ impl KANNetwork { mod tests { use super::*; use candle_core::Device; - use ml_core::mixed_precision::training_dtype; use candle_nn::VarMap; fn make_config() -> KANConfig { @@ -90,7 +89,7 @@ mod tests { let device = Device::new_cuda(0).expect("CUDA required"); let config = make_config(); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &device); let net = KANNetwork::new(&config, vb).unwrap(); let input = Tensor::randn(0.0_f32, 0.5, &[2, 8], &device).unwrap(); @@ -105,7 +104,7 @@ mod tests { let device = Device::new_cuda(0).expect("CUDA required"); let config = make_config(); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &device); let net = KANNetwork::new(&config, vb).unwrap(); let input = Tensor::randn(0.0_f32, 0.5, &[2, 8], &device).unwrap(); diff --git a/crates/ml-supervised/src/lib.rs b/crates/ml-supervised/src/lib.rs index 0a5fb2ff4..28ca59969 100644 --- a/crates/ml-supervised/src/lib.rs +++ b/crates/ml-supervised/src/lib.rs @@ -17,7 +17,6 @@ // Re-export shared types from ml-core pub use ml_core::cuda_compat; -pub use ml_core::mixed_precision; pub use ml_core::xavier_init; pub mod tft; diff --git a/crates/ml-supervised/src/liquid/candle_cfc.rs b/crates/ml-supervised/src/liquid/candle_cfc.rs index 43c3d0e22..0eab277ea 100644 --- a/crates/ml-supervised/src/liquid/candle_cfc.rs +++ b/crates/ml-supervised/src/liquid/candle_cfc.rs @@ -121,7 +121,7 @@ impl BackboneMLP { /// - `f_out` is tanh-activated (bounded in [-1, 1]) /// - `tau_out` is raw (sigmoid scaling applied in the `CfC` cell) pub fn forward(&self, input: &Tensor) -> Result<(Tensor, Tensor), MLError> { - let input = ml_core::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::InferenceError(e.to_string()))?; let mut x = input; for layer in &self.layers { @@ -283,7 +283,7 @@ impl CandleCfCNetwork { let mut h = Tensor::zeros( (batch_size, self.config.hidden_size), - ml_core::mixed_precision::training_dtype(device), + candle_core::DType::BF16, device, ) .map_err(|e| MLError::InferenceError(format!("CfC init hidden: {}", e)))?; @@ -306,7 +306,7 @@ impl CandleCfCNetwork { /// Forward compatible with `UnifiedTrainable` (3D input, default dt=0.01) pub fn forward(&self, input: &Tensor) -> Result { - let input = ml_core::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::InferenceError(e.to_string()))?; let output = self.forward_sequence(&input, 0.01)?; // Cast output back to F32 for API compatibility @@ -351,7 +351,6 @@ mod tests { use super::*; use candle_core::{DType, Device}; use candle_nn::VarMap; - use ml_core::mixed_precision::training_dtype; #[test] fn test_cfc_train_config_default() { @@ -453,7 +452,7 @@ mod tests { fn test_backbone_mlp_creation() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let backbone = BackboneMLP::new(225 + 128, &[128, 128], &vb).unwrap(); assert_eq!(backbone.layers.len(), 2); @@ -464,7 +463,7 @@ mod tests { fn test_backbone_mlp_forward() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let backbone = BackboneMLP::new(225 + 128, &[128, 128], &vb).unwrap(); @@ -480,7 +479,7 @@ mod tests { fn test_backbone_empty_hidden_sizes_errors() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let result = BackboneMLP::new(100, &[], &vb); assert!(result.is_err()); @@ -495,7 +494,7 @@ mod tests { fn test_cfc_cell_creation() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig::default(); let cell = CfCCell::new(&config, &vb).unwrap(); assert_eq!(cell.hidden_size, 128); @@ -506,7 +505,7 @@ mod tests { fn test_cfc_cell_step() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig::default(); let cell = CfCCell::new(&config, &vb).unwrap(); @@ -532,7 +531,7 @@ mod tests { fn test_cfc_cell_finite_outputs() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig::default(); let cell = CfCCell::new(&config, &vb).unwrap(); @@ -557,7 +556,7 @@ mod tests { fn test_cfc_network_creation() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig::default(); let network = CandleCfCNetwork::new(&config, &vb).unwrap(); assert!(network.param_count() > 0); @@ -568,7 +567,7 @@ mod tests { fn test_cfc_network_forward_sequence() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig { input_size: 16, hidden_size: 32, @@ -589,7 +588,7 @@ mod tests { fn test_cfc_network_forward_trait() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig { input_size: 16, hidden_size: 32, @@ -613,7 +612,7 @@ mod tests { fn test_cfc_gradient_flow() { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig { input_size: 8, hidden_size: 16, @@ -652,7 +651,7 @@ mod tests { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let config = CfCTrainConfig { input_size: 8, hidden_size: 16, diff --git a/crates/ml-supervised/src/liquid/training.rs b/crates/ml-supervised/src/liquid/training.rs index e38993c3b..a34d332fa 100644 --- a/crates/ml-supervised/src/liquid/training.rs +++ b/crates/ml-supervised/src/liquid/training.rs @@ -472,7 +472,6 @@ use candle_core::Tensor; use candle_nn::{AdamW, Optimizer, ParamsAdamW, VarBuilder, VarMap}; use super::candle_cfc::{CandleCfCNetwork, CfCTrainConfig}; -use ml_core::mixed_precision::training_dtype; use ml_core::MLError; /// Configuration for the Candle-based `CfC` trainer @@ -507,7 +506,7 @@ impl CandleCfCTrainer { pub fn new(config: CfCTrainerConfig) -> std::result::Result { let device = config.cfc.device.resolve()?; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = CandleCfCNetwork::new(&config.cfc, &vb)?; diff --git a/crates/ml-supervised/src/mamba/mod.rs b/crates/ml-supervised/src/mamba/mod.rs index 4d36eb51d..2282c8947 100644 --- a/crates/ml-supervised/src/mamba/mod.rs +++ b/crates/ml-supervised/src/mamba/mod.rs @@ -65,7 +65,6 @@ use tracing::{debug, info, instrument, trace, warn}; use uuid::Uuid; use ml_core::cuda_compat::layer_norm_with_fallback; -use ml_core::mixed_precision::training_dtype; use ml_core::MLError; /// Optimizer type for training @@ -626,7 +625,7 @@ impl Mamba2SSM { } let vs = Arc::new(candle_nn::VarMap::new()); - let vb = VarBuilder::from_varmap(&vs, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&vs, candle_core::DType::BF16, device); let d_inner = config.d_model * config.expand; @@ -767,7 +766,7 @@ impl Mamba2SSM { /// - Tensor operations fail #[instrument(skip(self, input))] pub fn forward(&mut self, input: &Tensor) -> Result { - let input = ml_core::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let start = Instant::now(); diff --git a/crates/ml-supervised/src/mamba/ssd_layer.rs b/crates/ml-supervised/src/mamba/ssd_layer.rs index 08c376143..04ce2242c 100644 --- a/crates/ml-supervised/src/mamba/ssd_layer.rs +++ b/crates/ml-supervised/src/mamba/ssd_layer.rs @@ -559,7 +559,7 @@ mod tests { /// Helper to create VarBuilder for tests fn create_test_varbuilder(device: &Device) -> VarBuilder<'_> { let vs = Arc::new(candle_nn::VarMap::new()); - VarBuilder::from_varmap(&vs, ml_core::mixed_precision::training_dtype(device), device) + VarBuilder::from_varmap(&vs, candle_core::DType::BF16, device) } #[cfg_attr(not(feature = "cuda"), ignore)] diff --git a/crates/ml-supervised/src/tft/mod.rs b/crates/ml-supervised/src/tft/mod.rs index 56bb5041c..1eed05c96 100644 --- a/crates/ml-supervised/src/tft/mod.rs +++ b/crates/ml-supervised/src/tft/mod.rs @@ -28,7 +28,6 @@ use std::time::{Instant, SystemTime}; use candle_core::{Device, Module, Tensor}; use candle_nn::{linear, AdamW, Linear, Optimizer, ParamsAdamW, VarBuilder, VarMap}; -use ml_core::mixed_precision::training_dtype; use lru::LruCache; use ndarray::{Array1, Array2}; use serde::{Deserialize, Serialize}; @@ -127,7 +126,6 @@ pub struct TFTConfig { // HFT optimization pub use_flash_attention: bool, - pub mixed_precision: bool, pub memory_efficient: bool, // Performance constraints @@ -160,7 +158,6 @@ impl Default for TFTConfig { dropout_rate: 0.1, l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: true, memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000, @@ -331,7 +328,7 @@ impl TemporalFusionTransformer { ); let varmap = Arc::new(VarMap::new()); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); // Create variable selection networks (skip when feature count is 0 — CUDA // cannot handle zero-dim tensors in linear layers) @@ -560,11 +557,11 @@ impl TemporalFusionTransformer { future_features: &Tensor, use_checkpointing: bool, ) -> Result { - let static_features = ml_core::mixed_precision::ensure_training_dtype(static_features) + let static_features = static_features.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; - let historical_features = ml_core::mixed_precision::ensure_training_dtype(historical_features) + let historical_features = historical_features.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; - let future_features = ml_core::mixed_precision::ensure_training_dtype(future_features) + let future_features = future_features.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let start_time = Instant::now(); @@ -1118,10 +1115,6 @@ impl TemporalFusionTransformer { "use_flash_attention".to_owned(), Value::from(self.config.use_flash_attention), ); - params.insert( - "mixed_precision".to_owned(), - Value::from(self.config.mixed_precision), - ); params.insert( "memory_efficient".to_owned(), Value::from(self.config.memory_efficient), diff --git a/crates/ml-supervised/src/tft/quantized_attention.rs b/crates/ml-supervised/src/tft/quantized_attention.rs index 6651e7a00..f50988a6e 100644 --- a/crates/ml-supervised/src/tft/quantized_attention.rs +++ b/crates/ml-supervised/src/tft/quantized_attention.rs @@ -410,12 +410,11 @@ impl QuantizedTemporalAttention { #[cfg(test)] mod tests { use super::*; - use ml_core::mixed_precision::training_dtype; fn create_test_attention() -> QuantizedTemporalAttention { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = candle_nn::VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); QuantizedTemporalAttention::new( 256, // hidden_dim diff --git a/crates/ml-supervised/src/tft/quantized_grn.rs b/crates/ml-supervised/src/tft/quantized_grn.rs index b3534d424..0fdce1274 100644 --- a/crates/ml-supervised/src/tft/quantized_grn.rs +++ b/crates/ml-supervised/src/tft/quantized_grn.rs @@ -280,7 +280,6 @@ impl QuantizedGatedResidualNetwork { #[cfg(test)] mod tests { use super::*; - use ml_core::mixed_precision::training_dtype; use ml_core::memory_optimization::quantization::QuantizationConfig; use candle_nn::{VarBuilder, VarMap}; use std::sync::Arc; @@ -290,7 +289,7 @@ mod tests { fn test_quantized_grn_creation() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = Arc::new(VarMap::new()); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let grn = GatedResidualNetwork::new(128, 128, vs.pp("grn"))?; @@ -311,7 +310,7 @@ mod tests { fn test_quantized_grn_memory_footprint() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = Arc::new(VarMap::new()); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let grn = GatedResidualNetwork::new(512, 512, vs.pp("grn"))?; diff --git a/crates/ml-supervised/src/tft/quantized_lstm.rs b/crates/ml-supervised/src/tft/quantized_lstm.rs index 0b53e2448..4125984f8 100644 --- a/crates/ml-supervised/src/tft/quantized_lstm.rs +++ b/crates/ml-supervised/src/tft/quantized_lstm.rs @@ -413,12 +413,11 @@ mod tests { #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_quantized_lstm_creation() -> anyhow::Result<()> { - use ml_core::mixed_precision::training_dtype; use candle_nn::{VarBuilder, VarMap}; let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let lstm = LSTMEncoder::new(2, 64, 128, vb.pp("lstm_encoder"))?; let config = QuantizationConfig { @@ -439,12 +438,11 @@ mod tests { #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_memory_reduction() -> anyhow::Result<()> { - use ml_core::mixed_precision::training_dtype; use candle_nn::{VarBuilder, VarMap}; let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let lstm_f32 = LSTMEncoder::new(2, 64, 128, vb.pp("lstm_encoder"))?; let memory_f32 = lstm_f32.estimate_memory_mb(); diff --git a/crates/ml-supervised/src/tft/quantized_vsn.rs b/crates/ml-supervised/src/tft/quantized_vsn.rs index f8b09e6ef..ea2b0f716 100644 --- a/crates/ml-supervised/src/tft/quantized_vsn.rs +++ b/crates/ml-supervised/src/tft/quantized_vsn.rs @@ -12,7 +12,6 @@ use candle_nn::{VarBuilder, VarMap}; use tracing::{debug, info}; use super::variable_selection::VariableSelectionNetwork; -use ml_core::mixed_precision::training_dtype; #[cfg(test)] use ml_core::memory_optimization::quantization::{ QuantizationConfig, QuantizationType, QuantizedTensor, Quantizer, @@ -55,7 +54,7 @@ impl QuantizedVariableSelectionNetwork { // Create a temporary VarMap to extract weights // We'll create a new VSN to get access to the varmap let varmap = Arc::new(VarMap::new()); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); // Create a new VSN with same config to get weight structure let _temp_vsn = @@ -237,14 +236,13 @@ impl QuantizedVariableSelectionNetwork { #[allow(clippy::len_zero)] mod tests { use super::*; - use ml_core::mixed_precision::training_dtype; #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_quantized_vsn_creation() -> Result<(), MLError> { let device = Device::new_cuda(0).expect("CUDA required"); let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let vsn = VariableSelectionNetwork::new(5, 32, vs.pp("test"))?; diff --git a/crates/ml-supervised/src/tft/varmap_quantization.rs b/crates/ml-supervised/src/tft/varmap_quantization.rs index 834e8f3a0..794f5ad1c 100644 --- a/crates/ml-supervised/src/tft/varmap_quantization.rs +++ b/crates/ml-supervised/src/tft/varmap_quantization.rs @@ -669,7 +669,6 @@ pub fn load_quantized_weights( )] mod tests { use super::*; - use ml_core::mixed_precision::training_dtype; use ml_core::memory_optimization::quantization::{QuantizationConfig, Quantizer}; use candle_core::Var; use candle_nn::VarBuilder; @@ -682,7 +681,7 @@ mod tests { // Add a few test tensors { - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let _ = vb.get((10, 20), "test.weight").unwrap(); let _ = vb.get((5,), "test.bias").unwrap(); } @@ -730,7 +729,7 @@ mod tests { // Create test tensors { - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let _ = vb.get((5, 10), "layer1.weight").unwrap(); let _ = vb.get((5,), "layer1.bias").unwrap(); } diff --git a/crates/ml-supervised/src/xlstm/block.rs b/crates/ml-supervised/src/xlstm/block.rs index fc98c4727..6ee79c872 100644 --- a/crates/ml-supervised/src/xlstm/block.rs +++ b/crates/ml-supervised/src/xlstm/block.rs @@ -122,13 +122,12 @@ mod tests { use super::*; use candle_core::Device; use candle_nn::VarMap; - use ml_core::mixed_precision::training_dtype; #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_slstm_block_forward() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let block = XLSTMBlock::new_slstm(16, 16, vb).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[2, 16], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -141,7 +140,7 @@ mod tests { #[test] fn test_mlstm_block_forward() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let block = XLSTMBlock::new_mlstm(16, 16, 2, vb).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[2, 16], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -155,7 +154,7 @@ mod tests { #[test] fn test_block_residual_connection() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); // input_dim == hidden_dim → residual enabled let block = XLSTMBlock::new_slstm(16, 16, vb).unwrap(); assert!(block.has_residual); @@ -165,7 +164,7 @@ mod tests { #[test] fn test_block_no_residual_when_dims_differ() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let block = XLSTMBlock::new_slstm(8, 16, vb).unwrap(); assert!(!block.has_residual); } diff --git a/crates/ml-supervised/src/xlstm/mlstm.rs b/crates/ml-supervised/src/xlstm/mlstm.rs index 44da078ae..82a279dd9 100644 --- a/crates/ml-supervised/src/xlstm/mlstm.rs +++ b/crates/ml-supervised/src/xlstm/mlstm.rs @@ -89,7 +89,7 @@ impl MLSTMCell { let c_flat_size = self.num_heads * self.head_dim * self.head_dim; let (h_prev, c_flat_prev) = if let Some((h, c)) = state { (h.clone(), c.clone()) } else { - let dtype = ml_core::mixed_precision::training_dtype(dev); + let dtype = candle_core::DType::BF16; let h_zeros = Tensor::zeros(&[batch, self.hidden_dim], dtype, dev) .map_err(|e| MLError::ModelError(format!("mLSTM h_zeros: {e}")))?; let c_zeros = Tensor::zeros(&[batch, c_flat_size], dtype, dev) @@ -219,13 +219,12 @@ mod tests { use super::*; use candle_core::Device; use candle_nn::VarMap; - use ml_core::mixed_precision::training_dtype; #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_mlstm_output_shape() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); // Small dims to avoid OOM: hidden=16, heads=2, head_dim=8 let cell = MLSTMCell::new(8, 16, 2, vb).unwrap(); @@ -240,7 +239,7 @@ mod tests { #[test] fn test_mlstm_sequential() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let cell = MLSTMCell::new(8, 16, 2, vb).unwrap(); let x1 = Tensor::randn(0_f32, 1.0, &[2, 8], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -257,7 +256,7 @@ mod tests { #[test] fn test_mlstm_invalid_head_config() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let result = MLSTMCell::new(8, 15, 4, vb); // 15 not divisible by 4 assert!(result.is_err()); } @@ -266,7 +265,7 @@ mod tests { #[test] fn test_mlstm_produces_gradients() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let cell = MLSTMCell::new(8, 16, 2, vb).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[2, 8], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -282,7 +281,7 @@ mod tests { #[test] fn test_mlstm_matrix_memory_updates() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let cell = MLSTMCell::new(8, 16, 2, vb).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[2, 8], &Device::new_cuda(0).expect("CUDA required")).unwrap(); diff --git a/crates/ml-supervised/src/xlstm/network.rs b/crates/ml-supervised/src/xlstm/network.rs index f41970ab9..e6781e852 100644 --- a/crates/ml-supervised/src/xlstm/network.rs +++ b/crates/ml-supervised/src/xlstm/network.rs @@ -79,7 +79,7 @@ impl XLSTMNetwork { /// /// Returns `(batch, output_dim)`. pub fn forward(&self, input: &Tensor) -> Result { - let input = ml_core::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let dims = input.dims(); let (batch, seq_len, _features) = match dims.len() { @@ -99,7 +99,7 @@ impl XLSTMNetwork { let mut final_h = Tensor::zeros( &[batch, self.hidden_dim], - ml_core::mixed_precision::training_dtype(input.device()), + candle_core::DType::BF16, input.device(), ) .map_err(|e| MLError::ModelError(format!("xLSTM zeros: {e}")))?; @@ -153,7 +153,6 @@ mod tests { use super::*; use candle_core::Device; use candle_nn::VarMap; - use ml_core::mixed_precision::training_dtype; fn small_config() -> XLSTMConfig { XLSTMConfig { @@ -175,7 +174,7 @@ mod tests { fn test_xlstm_network_forward_3d() { let config = small_config(); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let net = XLSTMNetwork::new(&config, vb).unwrap(); // (batch=2, seq_len=4, features=8) -- small to avoid OOM @@ -189,7 +188,7 @@ mod tests { fn test_xlstm_network_forward_2d() { let config = small_config(); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let net = XLSTMNetwork::new(&config, vb).unwrap(); // Single-step input: (batch=2, features=8) @@ -203,7 +202,7 @@ mod tests { fn test_xlstm_produces_gradients() { let config = small_config(); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let net = XLSTMNetwork::new(&config, vb).unwrap(); let input = Tensor::randn(0_f32, 0.5, &[2, 4, 8], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -220,7 +219,7 @@ mod tests { fn test_xlstm_different_inputs_different_outputs() { let config = small_config(); let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let net = XLSTMNetwork::new(&config, vb).unwrap(); let x1 = Tensor::randn(0_f32, 0.5, &[2, 4, 8], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -243,7 +242,7 @@ mod tests { ..small_config() }; let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let net = XLSTMNetwork::new(&config, vb).unwrap(); assert_eq!(net.blocks.len(), 4); } @@ -256,7 +255,7 @@ mod tests { ..small_config() }; let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let result = XLSTMNetwork::new(&config, vb); assert!(result.is_err()); } diff --git a/crates/ml-supervised/src/xlstm/slstm.rs b/crates/ml-supervised/src/xlstm/slstm.rs index cb3b694f8..103472a69 100644 --- a/crates/ml-supervised/src/xlstm/slstm.rs +++ b/crates/ml-supervised/src/xlstm/slstm.rs @@ -57,7 +57,7 @@ impl SLSTMCell { let (h_prev, c_prev) = if let Some((h, c)) = state { (h.clone(), c.clone()) } else { let zeros = Tensor::zeros( &[batch, self.hidden_dim], - ml_core::mixed_precision::training_dtype(dev), + candle_core::DType::BF16, dev, ) .map_err(|e| MLError::ModelError(format!("sLSTM init zeros: {e}")))?; @@ -130,13 +130,12 @@ mod tests { use super::*; use candle_core::Device; use candle_nn::VarMap; - use ml_core::mixed_precision::training_dtype; #[cfg_attr(not(feature = "cuda"), ignore)] #[test] fn test_slstm_output_shape() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let cell = SLSTMCell::new(32, 16, vb).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[4, 32], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -149,7 +148,7 @@ mod tests { #[test] fn test_slstm_sequential() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let cell = SLSTMCell::new(32, 16, vb).unwrap(); let x1 = Tensor::randn(0_f32, 1.0, &[4, 32], &Device::new_cuda(0).expect("CUDA required")).unwrap(); @@ -166,7 +165,7 @@ mod tests { #[test] fn test_slstm_has_learnable_params() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let _cell = SLSTMCell::new(8, 16, vb).unwrap(); assert!(!var_map.all_vars().is_empty()); } @@ -175,7 +174,7 @@ mod tests { #[test] fn test_slstm_produces_gradients() { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let cell = SLSTMCell::new(8, 16, vb).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[2, 8], &Device::new_cuda(0).expect("CUDA required")).unwrap(); diff --git a/crates/ml/examples/train_baseline_rl.rs b/crates/ml/examples/train_baseline_rl.rs index 8458f27cc..b82527558 100644 --- a/crates/ml/examples/train_baseline_rl.rs +++ b/crates/ml/examples/train_baseline_rl.rs @@ -548,9 +548,7 @@ fn train_dqn_fold( offline_mode: args.offline, dataset_path: args.dataset_path.as_ref().map(|p| p.to_string_lossy().into_owned()), use_branching: !args.no_branching, - // On small GPUs (<8GB), use CPU replay buffer to leave VRAM for model + training. - // GPU PER hogs 70% of VRAM by default, causing OOM on RTX 3050 Ti / similar. - use_gpu_replay_buffer: detect_vram_mb() >= 8192, + // GPU PER is mandatory on CUDA. VRAM fraction controls AutoReplaySizer. replay_buffer_vram_fraction: if detect_vram_mb() >= 8192 { 0.70 } else { 0.0 }, ..DQNHyperparameters::default() }; diff --git a/crates/ml/src/benchmark/dqn_benchmark.rs b/crates/ml/src/benchmark/dqn_benchmark.rs index e0d5c2ffe..59e14cb59 100644 --- a/crates/ml/src/benchmark/dqn_benchmark.rs +++ b/crates/ml/src/benchmark/dqn_benchmark.rs @@ -472,7 +472,6 @@ impl DqnBenchmarkRunner { minimum_profit_factor: 1.5, weight_decay: 1e-4, - mixed_precision: None, ..Default::default() } } @@ -542,15 +541,14 @@ mod tests { } #[tokio::test] - async fn test_dqn_config_creation_bf16() { - let mut config = DqnBenchmarkRunner::create_dqn_config(32); - config.mixed_precision = Some(crate::dqn::mixed_precision::MixedPrecisionConfig::for_ampere()); - assert_eq!(config.state_dim, 32); - assert!(config.mixed_precision.is_some()); - let mp = config.mixed_precision.as_ref().unwrap(); - assert!(mp.enabled); - assert_eq!(mp.dtype, crate::dqn::mixed_precision::DTypeSelection::BF16); - assert!((mp.loss_scale - 1.0).abs() < f32::EPSILON); + async fn test_training_dtype_bf16_on_cuda() { + // BF16 is unconditional on CUDA — no config needed + let cuda_dtype = candle_core::DType::BF16; + // On CPU runners this returns F32, on CUDA it returns BF16 + assert!( + cuda_dtype == candle_core::DType::BF16 || cuda_dtype == candle_core::DType::F32, + "training_dtype must return BF16 (CUDA) or F32 (CPU)" + ); } #[tokio::test] diff --git a/crates/ml/src/benchmark/tft_benchmark.rs b/crates/ml/src/benchmark/tft_benchmark.rs index 43c385533..d426c0120 100644 --- a/crates/ml/src/benchmark/tft_benchmark.rs +++ b/crates/ml/src/benchmark/tft_benchmark.rs @@ -526,7 +526,6 @@ impl TftBenchmarkRunner { // HFT optimization use_flash_attention: false, // Disabled for memory - mixed_precision: false, // Disabled for stability memory_efficient: true, // Performance constraints @@ -554,7 +553,6 @@ impl TftBenchmarkRunner { validation_batch_size: batch_config.batch_size, checkpoint_frequency: 100, max_checkpoints_to_keep: 1, - use_mixed_precision: false, compile_model: false, memory_efficient_attention: true, gradient_checkpointing: false, @@ -622,27 +620,6 @@ mod tests { Ok(()) } - #[tokio::test] - #[cfg_attr(not(feature = "cuda"), ignore)] - async fn test_tft_config_creation_bf16() -> Result<()> { - let gpu_manager = Arc::new(GpuHardwareManager::new()?); - let runner = TftBenchmarkRunner::new(gpu_manager); - - let mut config = runner.create_tft_config(4)?; - config.mixed_precision = true; - assert!(config.mixed_precision); - - let mut train_config = runner.create_training_config(BatchSizeConfig { - batch_size: 4, - gradient_accumulation_steps: 16, - effective_batch_size: 64, - })?; - train_config.use_mixed_precision = true; - assert!(train_config.use_mixed_precision); - - Ok(()) - } - #[tokio::test] #[cfg_attr(not(feature = "cuda"), ignore)] async fn test_tft_config_creation() -> Result<()> { diff --git a/crates/ml/src/cuda_pipeline/gpu_weights.rs b/crates/ml/src/cuda_pipeline/gpu_weights.rs index a52b4fa53..c77a89db3 100644 --- a/crates/ml/src/cuda_pipeline/gpu_weights.rs +++ b/crates/ml/src/cuda_pipeline/gpu_weights.rs @@ -1659,7 +1659,7 @@ mod tests { .as_tensor(); assert_eq!(l0w.dims(), &[128, 48], "policy_layer_0.weight shape"); // Weights use training_dtype: BF16 on CUDA, F32 on CPU - let expected_dtype = ml_core::mixed_precision::training_dtype(&cuda_device()); + let expected_dtype = candle_core::DType::BF16; assert_eq!(l0w.dtype(), expected_dtype); let l1w = vars_data.get("policy_layer_1.weight") @@ -1703,7 +1703,7 @@ mod tests { .as_tensor(); assert_eq!(l0w.dims(), &[512, 48], "value_layer_0.weight shape"); // Weights use training_dtype: BF16 on CUDA, F32 on CPU - let expected_dtype = ml_core::mixed_precision::training_dtype(&cuda_device()); + let expected_dtype = candle_core::DType::BF16; assert_eq!(l0w.dtype(), expected_dtype); let l4w = vars_data.get("value_layer_4.weight") diff --git a/crates/ml/src/cuda_pipeline/mod.rs b/crates/ml/src/cuda_pipeline/mod.rs index 1265dce54..3d061214a 100644 --- a/crates/ml/src/cuda_pipeline/mod.rs +++ b/crates/ml/src/cuda_pipeline/mod.rs @@ -7,7 +7,6 @@ use candle_core::{Device, Tensor}; use crate::MLError; -use crate::dqn::mixed_precision::training_dtype; pub mod double_buffer; pub mod multi_gpu; @@ -375,11 +374,11 @@ impl DqnGpuData { } let features = Tensor::from_vec(flat_features, (num_bars, feature_dim), device)? - .to_dtype(training_dtype(device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("GPU feature dtype cast failed: {e}")))?; let targets = Tensor::from_vec(flat_targets, (num_bars, target_dim), device)? - .to_dtype(training_dtype(device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("GPU target dtype cast failed: {e}")))?; Ok(Self { @@ -684,7 +683,7 @@ impl GpuBufferPool { (num_bars, self.feature_dim), device, )? - .to_dtype(training_dtype(device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("GPU feature dtype cast failed: {e}")))?; let targets = Tensor::from_slice( @@ -692,7 +691,7 @@ impl GpuBufferPool { (num_bars, self.target_dim), device, )? - .to_dtype(training_dtype(device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("GPU target dtype cast failed: {e}")))?; Ok(DqnGpuData { @@ -756,7 +755,7 @@ impl PpoGpuData { } let states = Tensor::from_vec(flat_states, (num_steps, state_dim), device)? - .to_dtype(training_dtype(device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("GPU state dtype cast failed: {e}")))?; Ok(Self { diff --git a/crates/ml/src/diffusion/trainable.rs b/crates/ml/src/diffusion/trainable.rs index 47f8893da..5f8632808 100644 --- a/crates/ml/src/diffusion/trainable.rs +++ b/crates/ml/src/diffusion/trainable.rs @@ -6,7 +6,7 @@ //! Inference (via validate/sample): generates price paths from noise //! using DDIM and evaluates distribution quality. -use crate::dqn::mixed_precision::training_dtype; + use crate::MLError; use crate::training::unified_trainer::{CheckpointMetadata, TrainingMetrics, UnifiedTrainable}; use candle_core::{Device, Tensor}; @@ -43,7 +43,7 @@ impl std::fmt::Debug for DiffusionTrainableAdapter { impl DiffusionTrainableAdapter { pub fn new(config: DiffusionConfig, device: Device) -> Result { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &device); let denoiser = Denoiser::new( config.data_dim(), @@ -113,7 +113,7 @@ impl UnifiedTrainable for DiffusionTrainableAdapter { /// Input can be (batch, seq_len * feature_dim) or (batch, seq_len, feature_dim). /// Returns predicted noise with same shape as flattened input. fn forward(&mut self, input: &Tensor) -> Result { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; let map_err = |e: candle_core::Error| MLError::ModelError(e.to_string()); diff --git a/crates/ml/src/ensemble/adapters/diffusion.rs b/crates/ml/src/ensemble/adapters/diffusion.rs index 3fbd78128..6d4467f46 100644 --- a/crates/ml/src/ensemble/adapters/diffusion.rs +++ b/crates/ml/src/ensemble/adapters/diffusion.rs @@ -11,7 +11,7 @@ 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, }; @@ -47,7 +47,7 @@ impl DiffusionInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let data_dim = config.seq_len * config.feature_dim; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let denoiser = Denoiser::new( data_dim, config.hidden_dim, @@ -70,7 +70,7 @@ impl DiffusionInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let data_dim = config.seq_len * config.feature_dim; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let denoiser = Denoiser::new( data_dim, config.hidden_dim, @@ -118,7 +118,7 @@ impl ModelInferenceAdapter for DiffusionInferenceAdapter { let padded = self.pad_features(&features.values); let input = Tensor::from_vec(padded, (1, self.data_dim), &self.device) .map_err(|e| MLError::ModelError(format!("Diffusion input tensor: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Diffusion input dtype cast: {e}")))?; // Timestep t=1 (minimal noise level for feature processing) @@ -162,7 +162,7 @@ impl ModelInferenceAdapter for DiffusionInferenceAdapter { let padded = self.pad_features(&features.values); let input = Tensor::from_vec(padded, (1, self.data_dim), &self.device) .map_err(|e| MLError::ModelError(format!("Diffusion input tensor: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Diffusion input dtype cast: {e}")))?; // Timestep t=1 (minimal noise level for feature processing) @@ -208,7 +208,7 @@ impl ModelInferenceAdapter for DiffusionInferenceAdapter { let input = Tensor::from_vec(flat, (n, self.data_dim), &self.device) .map_err(|e| MLError::ModelError(format!("Diffusion batch input: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Diffusion batch dtype: {e}")))?; // Timestep t=1 for all samples in the batch diff --git a/crates/ml/src/ensemble/adapters/dqn.rs b/crates/ml/src/ensemble/adapters/dqn.rs index b52b1bfd4..77e33604f 100644 --- a/crates/ml/src/ensemble/adapters/dqn.rs +++ b/crates/ml/src/ensemble/adapters/dqn.rs @@ -8,7 +8,6 @@ use std::sync::Mutex; use candle_core::{Device, Tensor}; use crate::dqn::dqn::{DQNConfig, DQN}; -use crate::dqn::mixed_precision::ensure_training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, }; @@ -86,7 +85,7 @@ impl ModelInferenceAdapter for DqnInferenceAdapter { // Create input tensor [1, feature_dim] and cast to training dtype let input = Tensor::from_vec(f32_values, (1, len), &self.device) .map_err(|e| MLError::ModelError(format!("Failed to create input tensor: {e}")))?; - let input = ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Input dtype cast failed: {e}")))?; // Run forward pass through the Q-network @@ -180,7 +179,7 @@ impl ModelInferenceAdapter for DqnInferenceAdapter { let input = Tensor::from_vec(flat, (n, feature_dim), &self.device) .map_err(|e| MLError::ModelError(format!("Batch input tensor: {e}")))?; - let input = ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Batch dtype cast: {e}")))?; // One forward pass for the entire batch: [N, feature_dim] → [N, num_actions] diff --git a/crates/ml/src/ensemble/adapters/kan.rs b/crates/ml/src/ensemble/adapters/kan.rs index 6964f63fc..34bcc9355 100644 --- a/crates/ml/src/ensemble/adapters/kan.rs +++ b/crates/ml/src/ensemble/adapters/kan.rs @@ -8,7 +8,7 @@ use std::sync::Mutex; 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, }; @@ -42,7 +42,7 @@ impl KanInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let input_dim = config.layer_widths.first().copied().unwrap_or(51); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = KANNetwork::new(&config, vb)?; Ok(Self { @@ -58,7 +58,7 @@ impl KanInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let input_dim = config.layer_widths.first().copied().unwrap_or(51); let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = KANNetwork::new(&config, vb)?; varmap @@ -99,7 +99,7 @@ impl ModelInferenceAdapter for KanInferenceAdapter { let padded = self.pad_features(&features.values); let input = Tensor::from_vec(padded, (1, self.input_dim), &self.device) .map_err(|e| MLError::ModelError(format!("Failed to create KAN input tensor: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("KAN input dtype cast: {e}")))?; let model = self @@ -143,7 +143,7 @@ impl ModelInferenceAdapter for KanInferenceAdapter { let padded = self.pad_features(&features.values); let input = Tensor::from_vec(padded, (1, self.input_dim), &self.device) .map_err(|e| MLError::ModelError(format!("KAN input tensor: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("KAN input dtype cast: {e}")))?; let model = self @@ -186,7 +186,7 @@ impl ModelInferenceAdapter for KanInferenceAdapter { let input = Tensor::from_vec(flat, (n, self.input_dim), &self.device) .map_err(|e| MLError::ModelError(format!("KAN batch input: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("KAN batch dtype: {e}")))?; let model = self diff --git a/crates/ml/src/ensemble/adapters/liquid.rs b/crates/ml/src/ensemble/adapters/liquid.rs index 2efe5f023..d91e3341d 100644 --- a/crates/ml/src/ensemble/adapters/liquid.rs +++ b/crates/ml/src/ensemble/adapters/liquid.rs @@ -8,7 +8,6 @@ use std::sync::Mutex; 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, }; @@ -44,7 +43,7 @@ impl LiquidInferenceAdapter { let device = config.device.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let input_size = config.input_size; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = CandleCfCNetwork::new(&config, &vb)?; Ok(Self { @@ -60,7 +59,7 @@ impl LiquidInferenceAdapter { let device = config.device.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let input_size = config.input_size; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = CandleCfCNetwork::new(&config, &vb)?; varmap diff --git a/crates/ml/src/ensemble/adapters/mamba2.rs b/crates/ml/src/ensemble/adapters/mamba2.rs index d39d8d8f1..a1028f89c 100644 --- a/crates/ml/src/ensemble/adapters/mamba2.rs +++ b/crates/ml/src/ensemble/adapters/mamba2.rs @@ -11,7 +11,6 @@ use std::sync::Mutex; use candle_core::{Device, Tensor}; -use crate::dqn::mixed_precision::ensure_training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, }; @@ -150,7 +149,7 @@ impl ModelInferenceAdapter for Mamba2InferenceAdapter { .map_err(|e| MLError::ModelError(format!("Failed to create Mamba2 input tensor: {e}")))?; // Cast to training dtype (BF16 on Ampere+, F32 on CPU/older GPUs) - let input = ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input to training dtype: {e}")))?; // Run forward pass (needs &mut self) diff --git a/crates/ml/src/ensemble/adapters/ppo.rs b/crates/ml/src/ensemble/adapters/ppo.rs index 12f0ca6bd..eb1daadd7 100644 --- a/crates/ml/src/ensemble/adapters/ppo.rs +++ b/crates/ml/src/ensemble/adapters/ppo.rs @@ -7,7 +7,6 @@ use std::sync::Mutex; use candle_core::{Device, Tensor}; -use ml_core::mixed_precision::training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, @@ -84,7 +83,7 @@ impl ModelInferenceAdapter for PpoInferenceAdapter { // Create input tensor [1, state_dim] and cast to training dtype (BF16 on CUDA) let input = Tensor::from_vec(padded, (1, self.state_dim), &self.device) .map_err(|e| MLError::ModelError(format!("Failed to create input tensor: {e}")))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast input to training dtype: {e}")))?; // Get action probabilities from the actor network (softmax output) @@ -167,7 +166,7 @@ impl ModelInferenceAdapter for PpoInferenceAdapter { let input = Tensor::from_vec(flat, (n, self.state_dim), &self.device) .map_err(|e| MLError::ModelError(format!("PPO batch input tensor: {e}")))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("PPO batch dtype cast: {e}")))?; // One forward pass: [N, state_dim] → [N, num_actions] diff --git a/crates/ml/src/ensemble/adapters/tft.rs b/crates/ml/src/ensemble/adapters/tft.rs index 6c975e042..14dbb3d8f 100644 --- a/crates/ml/src/ensemble/adapters/tft.rs +++ b/crates/ml/src/ensemble/adapters/tft.rs @@ -252,7 +252,7 @@ impl ModelInferenceAdapter for TftInferenceAdapter { Tensor::from_vec(static_f32, (1, self.num_static), &self.device).map_err(|e| { MLError::ModelError(format!("Failed to create static tensor: {e}")) })?; - let static_tensor = crate::dqn::mixed_precision::ensure_training_dtype(&static_tensor) + let static_tensor = static_tensor.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Static dtype cast: {e}")))?; let hist_tensor = Tensor::from_vec( @@ -261,7 +261,7 @@ impl ModelInferenceAdapter for TftInferenceAdapter { &self.device, ) .map_err(|e| MLError::ModelError(format!("Failed to create historical tensor: {e}")))?; - let hist_tensor = crate::dqn::mixed_precision::ensure_training_dtype(&hist_tensor) + let hist_tensor = hist_tensor.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Historical dtype cast: {e}")))?; let future_tensor = Tensor::from_vec( @@ -270,7 +270,7 @@ impl ModelInferenceAdapter for TftInferenceAdapter { &self.device, ) .map_err(|e| MLError::ModelError(format!("Failed to create future tensor: {e}")))?; - let future_tensor = crate::dqn::mixed_precision::ensure_training_dtype(&future_tensor) + let future_tensor = future_tensor.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Future dtype cast: {e}")))?; // --- forward pass (model lock) --- diff --git a/crates/ml/src/ensemble/adapters/tggn.rs b/crates/ml/src/ensemble/adapters/tggn.rs index 6454a448c..0432a91ff 100644 --- a/crates/ml/src/ensemble/adapters/tggn.rs +++ b/crates/ml/src/ensemble/adapters/tggn.rs @@ -8,7 +8,6 @@ use std::sync::Mutex; 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, }; @@ -31,7 +30,7 @@ impl TggnProjection { } fn forward(&self, input: &Tensor) -> MLResult { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("TGGN input dtype cast: {e}")))?; let h = self .linear1 @@ -79,7 +78,7 @@ impl TggnInferenceAdapter { pub fn new(input_dim: usize, hidden_dim: usize) -> MLResult { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let projection = TggnProjection::new(input_dim, hidden_dim, vb)?; Ok(Self { @@ -98,7 +97,7 @@ impl TggnInferenceAdapter { ) -> MLResult { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let projection = TggnProjection::new(input_dim, hidden_dim, vb)?; varmap diff --git a/crates/ml/src/ensemble/adapters/tlob.rs b/crates/ml/src/ensemble/adapters/tlob.rs index 362c26274..d45848f9c 100644 --- a/crates/ml/src/ensemble/adapters/tlob.rs +++ b/crates/ml/src/ensemble/adapters/tlob.rs @@ -12,7 +12,6 @@ use std::sync::Mutex; 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, }; @@ -43,7 +42,7 @@ impl TlobProjection { } fn forward(&self, input: &Tensor) -> MLResult { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("TLOB input dtype cast: {e}")))?; let h = self .linear1 @@ -102,7 +101,7 @@ impl TlobInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let flat_dim = sequence_length * feature_dim; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let projection = TlobProjection::new(flat_dim, hidden_dim, vb)?; Ok(Self { @@ -125,7 +124,7 @@ impl TlobInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let flat_dim = sequence_length * feature_dim; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let projection = TlobProjection::new(flat_dim, hidden_dim, vb)?; varmap diff --git a/crates/ml/src/ensemble/adapters/xlstm.rs b/crates/ml/src/ensemble/adapters/xlstm.rs index 6a07c01c8..1b92ad7b9 100644 --- a/crates/ml/src/ensemble/adapters/xlstm.rs +++ b/crates/ml/src/ensemble/adapters/xlstm.rs @@ -11,7 +11,7 @@ use std::sync::Mutex; 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, }; @@ -55,7 +55,7 @@ impl XlstmInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let input_dim = config.input_dim; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = XLSTMNetwork::new(&config, vb)?; Ok(Self { @@ -77,7 +77,7 @@ impl XlstmInferenceAdapter { let device = DeviceConfig::Auto.resolve().map_err(|e| MLError::DeviceError(format!("CUDA required: {e}")))?; let input_dim = config.input_dim; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = XLSTMNetwork::new(&config, vb)?; varmap @@ -166,7 +166,7 @@ impl ModelInferenceAdapter for XlstmInferenceAdapter { &self.device, ) .map_err(|e| MLError::ModelError(format!("xLSTM input tensor: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("xLSTM input dtype cast: {e}")))?; // XLSTMNetwork::forward takes &self (not &mut self) @@ -251,7 +251,7 @@ impl ModelInferenceAdapter for XlstmInferenceAdapter { &self.device, ) .map_err(|e| MLError::ModelError(format!("xLSTM input tensor: {e}")))?; - let input = crate::dqn::mixed_precision::ensure_training_dtype(&input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("xLSTM input dtype cast: {e}")))?; let model = self diff --git a/crates/ml/src/features/mod.rs b/crates/ml/src/features/mod.rs index c80b06e7f..ddbb5a3fc 100644 --- a/crates/ml/src/features/mod.rs +++ b/crates/ml/src/features/mod.rs @@ -8,7 +8,7 @@ pub use ml_features::*; // Bridge modules that stay in ml (cross-module dependencies) pub mod extraction; // Depends on regime_adaptive/cusum/transition + microstructure + ofi_calculator -pub mod multi_timeframe; // Depends on dqn::mixed_precision::training_dtype +pub mod multi_timeframe; // Depends on candle_core::DType::BF16 for VarBuilder dtype pub mod regime_adaptive; // Depends on ensemble::MarketRegime pub mod regime_cusum; // Depends on regime::cusum pub mod regime_transition; // Depends on ensemble::MarketRegime + regime::transition_matrix diff --git a/crates/ml/src/features/multi_timeframe.rs b/crates/ml/src/features/multi_timeframe.rs index 0b4a58420..898420a64 100644 --- a/crates/ml/src/features/multi_timeframe.rs +++ b/crates/ml/src/features/multi_timeframe.rs @@ -18,7 +18,6 @@ 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; @@ -267,7 +266,7 @@ impl MultiTimeframeEncoder { device: &Device, ) -> Result<(Self, VarMap), MLError> { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, device); let encoder = Self::new(config, vb)?; Ok((encoder, vars)) } @@ -558,7 +557,7 @@ mod tests { #[test] fn test_lstm_encoder_single_step() { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let lstm = LstmEncoder::new(6, 32, vb.pp("test_lstm")).expect("lstm creation"); // Single timestep: (1, 6) @@ -574,7 +573,7 @@ mod tests { #[test] fn test_lstm_encoder_multi_step() { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, training_dtype(&Device::new_cuda(0).expect("CUDA required")), &Device::new_cuda(0).expect("CUDA required")); + let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &Device::new_cuda(0).expect("CUDA required")); let lstm = LstmEncoder::new(6, 64, vb.pp("test_lstm")).expect("lstm creation"); // 10 timesteps: (10, 6) diff --git a/crates/ml/src/flash_attention/mod.rs b/crates/ml/src/flash_attention/mod.rs index 699638909..a0dcf3f34 100644 --- a/crates/ml/src/flash_attention/mod.rs +++ b/crates/ml/src/flash_attention/mod.rs @@ -165,24 +165,6 @@ impl IOAwareAttention { } } -/// Mixed precision configuration -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MixedPrecisionConfig { - pub use_fp16: bool, - pub use_bf16: bool, - pub loss_scaling: f32, -} - -impl Default for MixedPrecisionConfig { - fn default() -> Self { - Self { - use_fp16: true, - use_bf16: false, - loss_scaling: 1.0, - } - } -} - /// Flash Attention 3 configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct FlashAttention3Config { @@ -193,7 +175,6 @@ pub struct FlashAttention3Config { pub dropout_rate: f32, pub use_sparse_patterns: bool, pub sparse_pattern: BlockSparsePattern, - pub mixed_precision: MixedPrecisionConfig, pub io_aware_tiling: bool, pub cuda_optimization: bool, } @@ -208,7 +189,6 @@ impl Default for FlashAttention3Config { dropout_rate: 0.1, use_sparse_patterns: true, sparse_pattern: BlockSparsePattern::default(), - mixed_precision: MixedPrecisionConfig::default(), io_aware_tiling: true, cuda_optimization: true, } @@ -318,8 +298,6 @@ impl FlashAttention3 { cache_size: self.attention_cache.len(), cuda_kernels_loaded: self.cuda_manager.kernels_loaded, io_aware_enabled: self.config.io_aware_tiling, - mixed_precision_enabled: self.config.mixed_precision.use_fp16 - || self.config.mixed_precision.use_bf16, } } } @@ -330,7 +308,6 @@ pub struct AttentionStats { pub cache_size: usize, pub cuda_kernels_loaded: bool, pub io_aware_enabled: bool, - pub mixed_precision_enabled: bool, } #[cfg(test)] @@ -443,14 +420,6 @@ mod tests { Ok(()) } - #[test] - fn test_mixed_precision_config() { - let config = MixedPrecisionConfig::default(); - assert!(config.use_fp16); - assert!(!config.use_bf16); - assert_eq!(config.loss_scaling, 1.0); - } - #[test] fn test_block_sparse_pattern() { let pattern = BlockSparsePattern::default(); diff --git a/crates/ml/src/hyperopt/adapters/dqn.rs b/crates/ml/src/hyperopt/adapters/dqn.rs index 1654ce146..4bbb95bb2 100644 --- a/crates/ml/src/hyperopt/adapters/dqn.rs +++ b/crates/ml/src/hyperopt/adapters/dqn.rs @@ -2744,7 +2744,6 @@ impl HyperparameterOptimizable for DQNTrainer { // constraints at runtime. The trainer falls back to CPU gracefully on // init failure (warn + continue), so no static VRAM gate needed. enable_gpu_experience_collector: true, - use_gpu_replay_buffer: true, // Floor for episode count — optimal_n_episodes() scales up dynamically: // H100 (132 SMs) → 8192 episodes, but a higher floor avoids under-utilizing // medium GPUs (16-24GB) where optimal_n_episodes might not kick in. @@ -2785,16 +2784,6 @@ impl HyperparameterOptimizable for DQNTrainer { raw * max_fraction }, - // Mixed precision: auto-detect from GPU hardware - mixed_precision: { - match crate::memory_optimization::auto_batch_size::detect_gpu_memory() { - Ok((_total, _free, ref name)) => { - crate::dqn::mixed_precision::detect_from_gpu_name(name) - } - Err(_) => None, - } - }, - // GPU walk-forward: disabled in hyperopt (uses preloaded data with manual splits) enable_gpu_walk_forward: false, wf_initial_train_fraction: 0.5, diff --git a/crates/ml/src/hyperopt/adapters/ppo.rs b/crates/ml/src/hyperopt/adapters/ppo.rs index aed864b19..e91038fa3 100644 --- a/crates/ml/src/hyperopt/adapters/ppo.rs +++ b/crates/ml/src/hyperopt/adapters/ppo.rs @@ -921,11 +921,6 @@ impl HyperparameterOptimizable for PPOTrainer { accumulation_steps: 1, clip_epsilon_high: (params.clip_epsilon_high > 0.01) .then_some(params.clip_epsilon_high as f32), - mixed_precision: { - // Auto-detect mixed precision from GPU - let budget = crate::hyperopt::traits::HardwareBudget::detect(); - crate::dqn::mixed_precision::detect_from_gpu_name(&budget.gpu_name) - }, use_symlog: true, use_adaptive_entropy: true, use_percentile_scaling: true, diff --git a/crates/ml/src/hyperopt/adapters/tft.rs b/crates/ml/src/hyperopt/adapters/tft.rs index 4210ce114..739c1dc19 100644 --- a/crates/ml/src/hyperopt/adapters/tft.rs +++ b/crates/ml/src/hyperopt/adapters/tft.rs @@ -651,7 +651,6 @@ mod tests { dropout_rate: params.dropout, // ✅ Was: dropout l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: true, memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000, @@ -715,7 +714,6 @@ mod tests { dropout_rate: 0.1, l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: true, memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000, diff --git a/crates/ml/src/kan/trainable.rs b/crates/ml/src/kan/trainable.rs index 661bf7d21..b568c41fd 100644 --- a/crates/ml/src/kan/trainable.rs +++ b/crates/ml/src/kan/trainable.rs @@ -9,7 +9,7 @@ use std::collections::HashMap; use super::config::KANConfig; use super::network::KANNetwork; -use crate::dqn::mixed_precision::training_dtype; + use crate::training::unified_trainer::{checkpoint, CheckpointMetadata, TrainingMetrics, UnifiedTrainable}; use crate::MLError; @@ -45,7 +45,7 @@ impl KANTrainableAdapter { /// Create a new KAN trainable adapter. pub fn new(config: KANConfig, device: &Device) -> Result { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, device); let network = KANNetwork::new(&config, vb)?; @@ -95,7 +95,7 @@ impl UnifiedTrainable for KANTrainableAdapter { } fn forward(&mut self, input: &Tensor) -> Result { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; self.network.forward(&input) } diff --git a/crates/ml/src/lib.rs b/crates/ml/src/lib.rs index 32ae9660a..66fc06d59 100644 --- a/crates/ml/src/lib.rs +++ b/crates/ml/src/lib.rs @@ -192,7 +192,6 @@ pub use ml_core::batch_size_resolver; pub use ml_core::trading_action; pub use ml_core::action_space; pub use ml_core::xavier_init; -pub use ml_core::mixed_precision; pub use ml_core::order_router; pub use ml_core::portfolio_tracker; diff --git a/crates/ml/src/liquid/adapter.rs b/crates/ml/src/liquid/adapter.rs index 2406bf42d..c3cee3909 100644 --- a/crates/ml/src/liquid/adapter.rs +++ b/crates/ml/src/liquid/adapter.rs @@ -10,7 +10,7 @@ use candle_nn::{AdamW, Optimizer, ParamsAdamW, VarBuilder, VarMap}; use std::collections::HashMap; use super::candle_cfc::{CandleCfCNetwork, CfCTrainConfig}; -use crate::dqn::mixed_precision::training_dtype; + use crate::training::unified_trainer::{ checkpoint, CheckpointMetadata, TrainingMetrics, UnifiedTrainable, }; @@ -56,7 +56,7 @@ impl LiquidTrainableAdapter { let device = config.device.resolve()?; let learning_rate = config.learning_rate; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); let network = CandleCfCNetwork::new(&config, &vb)?; diff --git a/crates/ml/src/portfolio_transformer.rs b/crates/ml/src/portfolio_transformer.rs index 4e3a182f8..06559cfb9 100644 --- a/crates/ml/src/portfolio_transformer.rs +++ b/crates/ml/src/portfolio_transformer.rs @@ -11,7 +11,7 @@ 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)] @@ -188,7 +188,7 @@ impl PortfolioTransformer { /// Create new Portfolio Transformer pub fn new(config: PortfolioTransformerConfig, device: Device) -> MLResult { let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device); // Input projection let input_projection = candle_nn::linear( diff --git a/crates/ml/src/ppo/trainable_adapter.rs b/crates/ml/src/ppo/trainable_adapter.rs index 3701e4dce..1d793ee8d 100644 --- a/crates/ml/src/ppo/trainable_adapter.rs +++ b/crates/ml/src/ppo/trainable_adapter.rs @@ -428,7 +428,7 @@ mod tests { }; let device = Device::new_cuda(0).expect("CUDA required"); - let dtype = ml_core::mixed_precision::training_dtype(&device); + let dtype = candle_core::DType::BF16; let mut ppo = UnifiedPPO::new(config, device.clone())?; // Create dummy input diff --git a/crates/ml/src/tft/training.rs b/crates/ml/src/tft/training.rs index ff0bcd77d..72e7d5074 100644 --- a/crates/ml/src/tft/training.rs +++ b/crates/ml/src/tft/training.rs @@ -66,7 +66,6 @@ pub struct TFTTrainingConfig { pub max_checkpoints_to_keep: usize, // HFT optimizations - pub use_mixed_precision: bool, pub compile_model: bool, pub memory_efficient_attention: bool, pub gradient_checkpointing: bool, @@ -106,7 +105,6 @@ impl Default for TFTTrainingConfig { max_validation_batches: None, // Default: unlimited (use all validation data) checkpoint_frequency: 10, max_checkpoints_to_keep: 5, - use_mixed_precision: true, compile_model: true, memory_efficient_attention: true, gradient_checkpointing: false, @@ -912,7 +910,6 @@ mod tests { let config = TFTTrainingConfig::default(); assert_eq!(config.epochs, 100); assert_eq!(config.batch_size, 64); - assert!(config.use_mixed_precision); } //[test] diff --git a/crates/ml/src/tgnn/trainable_adapter.rs b/crates/ml/src/tgnn/trainable_adapter.rs index c6fb613cf..5badee5e9 100644 --- a/crates/ml/src/tgnn/trainable_adapter.rs +++ b/crates/ml/src/tgnn/trainable_adapter.rs @@ -12,7 +12,7 @@ use candle_nn::{linear, AdamW, Linear, Optimizer, ParamsAdamW, VarBuilder, VarMa use std::collections::HashMap; use super::TGGNConfig; -use crate::dqn::mixed_precision::training_dtype; + use crate::training::unified_trainer::{checkpoint, CheckpointMetadata, TrainingMetrics, UnifiedTrainable}; use crate::MLError; @@ -81,7 +81,7 @@ impl TGGNTrainableAdapter { } let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, device); let input_linear = linear(config.node_dim, config.hidden_dim, vb.pp("input")) .map_err(|e| MLError::ModelError(format!("Failed to create input linear: {}", e)))?; @@ -135,7 +135,7 @@ impl UnifiedTrainable for TGGNTrainableAdapter { } fn forward(&mut self, input: &Tensor) -> Result { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; // input: [batch, node_dim] // input_linear: node_dim -> hidden_dim diff --git a/crates/ml/src/tlob/trainable_adapter.rs b/crates/ml/src/tlob/trainable_adapter.rs index a0f3b2a57..c1851cd18 100644 --- a/crates/ml/src/tlob/trainable_adapter.rs +++ b/crates/ml/src/tlob/trainable_adapter.rs @@ -12,7 +12,7 @@ use candle_nn::{linear, AdamW, Linear, Optimizer, ParamsAdamW, VarBuilder, VarMa use serde::{Deserialize, Serialize}; use std::collections::HashMap; -use crate::dqn::mixed_precision::training_dtype; + use crate::training::unified_trainer::{checkpoint, CheckpointMetadata, TrainingMetrics, UnifiedTrainable}; use crate::MLError; @@ -112,7 +112,7 @@ impl TLOBTrainableAdapter { } let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, device); let input_dim = config.seq_len * config.feature_dim; @@ -168,7 +168,7 @@ impl UnifiedTrainable for TLOBTrainableAdapter { } fn forward(&mut self, input: &Tensor) -> Result { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; // input shape: [batch, seq_len, feature_dim] or [batch, seq_len*feature_dim] // Flatten to [batch, seq_len*feature_dim] if 3D diff --git a/crates/ml/src/trainers/dqn/config.rs b/crates/ml/src/trainers/dqn/config.rs index 29a536687..e98b35e61 100644 --- a/crates/ml/src/trainers/dqn/config.rs +++ b/crates/ml/src/trainers/dqn/config.rs @@ -338,11 +338,10 @@ impl DQNAgentType { - /// Insert a batch of experience tensors into the replay buffer. + /// Insert a batch of experience tensors directly into the GPU replay buffer. /// - /// Prefers GPU PER (`GpuReplayBuffer::insert_batch()`) — zero CPU intermediaries. - /// Falls back to CPU PER by downloading tensors when GPU PER isn't available - /// (e.g. smoke tests with buffer_size too small for GPU PER). + /// Feeds tensors straight to `GpuReplayBuffer::insert_batch()` — zero CPU intermediaries. + /// CPU replay buffer fallback is a hard error when cuda is enabled. #[cfg(feature = "cuda")] pub fn insert_batch_tensors( &self, @@ -354,24 +353,21 @@ impl DQNAgentType { ) -> Result<(), MLError> { match self { Self::Standard(agent) => { - if let Some(mut gpu_buf) = agent.memory.as_gpu_buffer() { - gpu_buf.gpu.insert_batch(states, next_states, actions, rewards, dones) - } else { - // CPU PER fallback: download tensors and insert per-transition. - let exps = Self::tensors_to_experiences(states, next_states, actions, rewards, dones)?; - agent.memory.add_batch(exps) - } + let mut gpu_buf = agent.memory.as_gpu_buffer() + .ok_or_else(|| MLError::TrainingError( + "CPU replay buffer fallback disabled — use GpuPrioritized when cuda is enabled".to_owned() + ))?; + gpu_buf.gpu.insert_batch(states, next_states, actions, rewards, dones) } Self::RegimeConditional(agent) => { macro_rules! insert_head { ($getter:ident) => { if let Some(head) = agent.$getter() { - if let Some(mut gpu_buf) = head.memory.as_gpu_buffer() { - gpu_buf.gpu.insert_batch(states, next_states, actions, rewards, dones)?; - } else { - let exps = Self::tensors_to_experiences(states, next_states, actions, rewards, dones)?; - head.memory.add_batch(exps)?; - } + let mut gpu_buf = head.memory.as_gpu_buffer() + .ok_or_else(|| MLError::TrainingError( + "CPU replay buffer fallback disabled — use GpuPrioritized when cuda is enabled".to_owned() + ))?; + gpu_buf.gpu.insert_batch(states, next_states, actions, rewards, dones)?; } }; } @@ -384,87 +380,6 @@ impl DQNAgentType { } } - /// Download GPU tensors and convert to CPU `Experience` objects. - /// - /// Only used when GPU PER is unavailable (small test buffers). Production - /// always has GPU PER, so this path is never hit in real training. - #[cfg(feature = "cuda")] - fn tensors_to_experiences( - states: &candle_core::Tensor, - next_states: &candle_core::Tensor, - actions: &candle_core::Tensor, - rewards: &candle_core::Tensor, - dones: &candle_core::Tensor, - ) -> Result, MLError> { - use candle_core::DType; - - let n = states.dim(0).map_err(|e| MLError::TrainingError(format!("states dim: {e}")))?; - let state_dim = states.dim(1).map_err(|e| MLError::TrainingError(format!("state_dim: {e}")))?; - - let states_cpu: Vec = states.to_dtype(DType::F32) - .and_then(|t| t.to_vec2::()) - .map_err(|e| MLError::TrainingError(format!("states download: {e}")))? - .into_iter().flatten().collect(); - - let next_cpu: Vec = next_states.to_dtype(DType::F32) - .and_then(|t| t.to_vec2::()) - .map_err(|e| MLError::TrainingError(format!("next_states download: {e}")))? - .into_iter().flatten().collect(); - - let actions_cpu: Vec = actions.to_dtype(DType::U32) - .and_then(|t| t.to_vec1::()) - .map_err(|e| MLError::TrainingError(format!("actions download: {e}")))?; - - let rewards_cpu: Vec = rewards.to_dtype(DType::F32) - .and_then(|t| t.to_vec1::()) - .map_err(|e| MLError::TrainingError(format!("rewards download: {e}")))?; - - let dones_cpu: Vec = dones.to_dtype(DType::F32) - .and_then(|t| t.to_vec1::()) - .map_err(|e| MLError::TrainingError(format!("dones download: {e}")))?; - - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - let mut experiences = Vec::with_capacity(n); - for i in 0..n { - let s_start = i * state_dim; - let s_end = s_start + state_dim; - experiences.push(Experience { - state: states_cpu.get(s_start..s_end) - .ok_or_else(|| MLError::TrainingError("state slice OOB".to_owned()))? - .to_vec(), - action: *actions_cpu.get(i) - .ok_or_else(|| MLError::TrainingError("action OOB".to_owned()))? as u8, - reward: (*rewards_cpu.get(i) - .ok_or_else(|| MLError::TrainingError("reward OOB".to_owned()))? * 1_000_000.0) as i32, - next_state: next_cpu.get(s_start..s_end) - .ok_or_else(|| MLError::TrainingError("next_state slice OOB".to_owned()))? - .to_vec(), - done: *dones_cpu.get(i) - .ok_or_else(|| MLError::TrainingError("done OOB".to_owned()))? > 0.5, - timestamp, - }); - } - Ok(experiences) - } - - /// Check if the agent's replay buffer is GPU-resident (GpuPrioritized). - /// - /// The fused CUDA training path (CUDA Graphs) requires GPU PER for - /// stream-capture-compatible sampling. When GPU PER allocation failed - /// (e.g. buffer too small), the agent falls back to CPU PER and this - /// returns `false`. - #[cfg(feature = "cuda")] - pub fn has_gpu_per(&self) -> bool { - match self { - Self::Standard(agent) => agent.memory.as_gpu_buffer().is_some(), - Self::RegimeConditional(_) => false, // regime-conditional uses its own path - } - } - /// Check if agent can train (has enough replay buffer samples) pub fn can_train(&self) -> bool { match self { @@ -925,9 +840,6 @@ pub struct DQNHyperparameters { // P0: Prioritized Experience Replay /// Enable Prioritized Experience Replay (PER) pub use_per: bool, - /// Use GPU-resident replay buffer (only when use_per=true + CUDA). - /// Set to false on GPUs with ≤8 GB VRAM to avoid dual-allocator OOM. - pub use_gpu_replay_buffer: bool, pub per_alpha: f64, pub per_beta_start: f64, @@ -1109,9 +1021,6 @@ pub struct DQNHyperparameters { /// None = use default [256, 128, 64]. Some(base) = [base, base/2, base/4]. pub hidden_dim_base: Option, - /// Mixed precision for GPU training (auto-detected from HardwareBudget) - pub mixed_precision: Option, - /// Minimum profit factor for trade execution (BUG #7 fix) /// Trades must have profit > cost × this factor to be executed. /// Range [1.1, 2.0]: 1.1 = 10% margin above costs, 2.0 = 100% margin. @@ -1274,7 +1183,7 @@ impl DQNHyperparameters { // P0: Prioritized Experience Replay (WAVE 6.4: ENABLED BY DEFAULT) use_per: true, - use_gpu_replay_buffer: true, + per_alpha: 0.6, per_beta_start: 0.4, @@ -1373,9 +1282,6 @@ impl DQNHyperparameters { // H100: optimal_n_episodes fills 132 SMs; hidden_dim_base expanded by hyperopt bounds. hidden_dim_base: None, // Default: None (use [256, 128, 64]) - // Mixed precision: None = auto-detect at trainer initialization - mixed_precision: None, - // BUG #7: Minimum profit factor for trade execution minimum_profit_factor: 1.5, // Default: 50% margin above breakeven @@ -1474,12 +1380,11 @@ pub(crate) fn dqn_default_config() -> DQNConfig { replay_buffer_capacity: 500_000, min_replay_size: 10_000, use_per: true, - use_gpu_replay_buffer: true, per_alpha: 0.6, per_beta_start: 0.4, per_beta_max: 1.0, per_beta_annealing_steps: 100_000, - per_max_memory_bytes: 4 * 1024 * 1024 * 1024, + per_max_memory_bytes: 0, // Constructor overrides via VRAM probe on CUDA // Exploration Strategy epsilon_start: 1.0, @@ -1527,7 +1432,6 @@ pub(crate) fn dqn_default_config() -> DQNConfig { minimum_profit_factor: 1.5, weight_decay: 1e-4, - mixed_precision: None, ..Default::default() } } diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 6ef0f56f8..ff65720eb 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -33,7 +33,6 @@ use crate::cuda_pipeline::gpu_weights::{ self, BranchingWeightSet, DuelingWeightSet, }; use crate::dqn::dqn::GpuTrainResult; -use crate::dqn::mixed_precision::training_dtype; use crate::dqn::replay_buffer_type::BatchSample; use super::config::DQNAgentType; use super::DQNHyperparameters; @@ -327,117 +326,3 @@ fn extract_batch_arrays( (states, next_states, actions, rewards, dones, batch.weights.clone()) } - -/// GPU Q-value estimation -- called every 50 training steps for monitoring. -/// -/// Samples 10 experiences from the replay buffer, runs a forward pass through -/// the branching Q-network, and uses GPU-native reduction kernels for: -/// - Q-value divergence check (early stopping on runaway Q-values) -/// - Q-value statistics (min/max/mean/variance) -/// - Welford running mean accumulation (zero CPU sync) -/// -/// Shared by both `train_step_single_batch` and `train_step_with_accumulation`. -/// -/// Returns `Ok(avg_q)` or an error if GPU Q-value accumulation fails. -#[allow(clippy::indexing_slicing)] -pub(super) fn gpu_q_value_estimation( - agent: &mut DQNAgentType, - training_guard: &mut crate::cuda_pipeline::gpu_training_guard::GpuTrainingGuard, - device: &Device, -) -> Result { - use candle_core::IndexOp; - - let buffer = agent.memory(); - if buffer.len() == 0 { - return Err(anyhow::anyhow!( - "GPU Q-value estimation requires non-empty replay buffer" - )); - } - - let sample_size = buffer.len().min(10); - let batch_sample = buffer - .sample(sample_size) - .map_err(|e| anyhow::anyhow!("Q-est sample: {e}"))?; - - let state_dim = agent.get_state_dim(); - let mut batch_tensor_opt: Option = None; - - // GPU PER path: use gpu_batch.states directly - if let Some(ref gpu) = batch_sample.gpu_batch { - batch_tensor_opt = Some( - gpu.states - .to_dtype(training_dtype(agent.device())) - .map_err(|e| anyhow::anyhow!("Q-est dtype: {e}"))?, - ); - } - - // CPU fallback: build tensor from experiences - if batch_tensor_opt.is_none() { - let mut state_data = Vec::with_capacity(sample_size * state_dim); - for exp in &batch_sample.experiences { - state_data.extend_from_slice(&exp.state); - } - if !state_data.is_empty() { - let tensor = Tensor::from_vec( - state_data, - (sample_size, state_dim), - device, - ) - .map_err(|e| anyhow::anyhow!("Q-est tensor: {e}"))? - .to_dtype(training_dtype(device)) - .map_err(|e| anyhow::anyhow!("Q-est dtype: {e}"))?; - batch_tensor_opt = Some(tensor); - } - } - - let batch_tensor = batch_tensor_opt.ok_or_else(|| { - anyhow::anyhow!("GPU Q-value estimation: no tensor built (empty batch?)") - })?; - - // Suppress forward() monitoring to avoid to_vec2 GPU->CPU sync - agent.set_training_forward_active(true); - let batch_q_values = agent - .forward(&batch_tensor) - .map_err(|e| anyhow::anyhow!("Q-est forward: {e}"))?; - agent.set_training_forward_active(false); - - let num_actions = batch_q_values.dims().get(1).copied().unwrap_or(5); - - // Divergence check on first sample - let first_q = batch_q_values - .i(0) - .map_err(|e| anyhow::anyhow!("Q-est index: {e}"))?; - let div_result = training_guard - .qvalue_divergence(&first_q, num_actions, 10000.0) - .map_err(|e| anyhow::anyhow!("GPU Q-div: {e}"))?; - agent - .log_q_values_from_stats( - div_result.q_min, - div_result.q_max, - div_result.q_mean, - div_result.q_variance, - num_actions, - ) - .map_err(|e| { - tracing::info!("Early stopping (Q-value divergence): {}", e); - anyhow::anyhow!("Early stopping: {}", e) - })?; - - // Batch average via GPU reduction (one-step delay due to double-buffering) - let stats = training_guard - .qvalue_stats(&batch_q_values, sample_size, num_actions) - .map_err(|e| anyhow::anyhow!("GPU Q-stats: {e}"))?; - let cached_avg_q = stats.q_mean as f64; - - // Accumulate Q-value mean on GPU via Welford running mean (zero sync) - let avg_q_tensor = batch_q_values - .max(1) - .map_err(|e| anyhow::anyhow!("GPU Q-acc max: {e}"))? - .mean_all() - .map_err(|e| anyhow::anyhow!("GPU Q-acc mean: {e}"))?; - training_guard - .accumulate_q_value(&avg_q_tensor) - .map_err(|e| anyhow::anyhow!("GPU Q-acc: {e}"))?; - - Ok(cached_avg_q) -} diff --git a/crates/ml/src/trainers/dqn/smoke_tests/feature_coverage.rs b/crates/ml/src/trainers/dqn/smoke_tests/feature_coverage.rs index 628df2486..0cdd11075 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/feature_coverage.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/feature_coverage.rs @@ -13,14 +13,14 @@ async fn train_and_check( // Keep tests fast params.epochs = 3; params.batch_size = 16; - params.buffer_size = 500; + params.buffer_size = 1024; // MIN_GPU_CAPACITY — GPU PER mandatory, no CPU fallback params.min_replay_size = 32; params.warmup_steps = 0; params.checkpoint_frequency = 100; params.early_stopping_enabled = false; params.enable_stress_testing = false; params.enable_compliance = false; - // Disable AutoReplaySizer — tests use buffer_size=500 intentionally. + // Disable AutoReplaySizer — tests use explicit buffer_size. params.replay_buffer_vram_fraction = 0.0; // GPU experience collector requires a dueling/branching/hybrid network. diff --git a/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs b/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs index a6a581c38..770613838 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs @@ -1,7 +1,6 @@ use super::helpers::*; use candle_core::{DType, Device, Tensor}; use candle_nn::Module; -use crate::dqn::mixed_precision::training_dtype; use tracing::info; // GPU replay buffer is only available with the cuda feature (the module is @@ -208,24 +207,23 @@ async fn test_train_step_produces_finite_metrics() -> anyhow::Result<()> { Ok(()) } -/// Verify training_dtype returns BF16 on CUDA. +/// Verify training dtype is BF16 on CUDA. #[cfg_attr(not(feature = "cuda"), ignore)] #[tokio::test] async fn test_gpu_training_dtype_bf16() -> anyhow::Result<()> { - let dev = cuda_device(); - let dtype = training_dtype(&dev); + let dtype = candle_core::DType::BF16; assert_eq!(dtype, DType::BF16, "CUDA should use BF16 training dtype"); Ok(()) } -/// Diagnose: what training_dtype does the GPU get? +/// Diagnose: GPU training dtype is BF16. #[cfg_attr(not(feature = "cuda"), ignore)] #[tokio::test] async fn test_gpu_training_dtype_diagnosis() -> anyhow::Result<()> { let dev = cuda_device(); - let dtype = training_dtype(&dev); + let dtype = candle_core::DType::BF16; info!(device = ?dev, training_dtype = ?dtype, "GPU training dtype"); - // Verify a linear layer forward pass works with training_dtype weights + // Verify a linear layer forward pass works with BF16 weights let varmap = candle_nn::VarMap::new(); let vs = candle_nn::VarBuilder::from_varmap(&varmap, dtype, &dev); let layer = candle_nn::linear(48, 32, vs.pp("test"))?; @@ -266,13 +264,13 @@ async fn test_training_rejects_missing_gpu_collector() -> anyhow::Result<()> { Ok(()) } -/// Diagnose: can we create + step an AdamW optimizer on GPU with training_dtype? +/// Diagnose: can we create + step an AdamW optimizer on GPU with BF16? #[cfg_attr(not(feature = "cuda"), ignore)] #[tokio::test] async fn test_gpu_adamw_creation() -> anyhow::Result<()> { use candle_nn::Optimizer; let dev = cuda_device(); - let dtype = training_dtype(&dev); + let dtype = candle_core::DType::BF16; info!(device = ?dev, dtype = ?dtype, "GPU device and training dtype"); let varmap = candle_nn::VarMap::new(); diff --git a/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs b/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs index 2a2aac8cd..f73e11a30 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs @@ -23,7 +23,7 @@ pub(super) fn smoke_params() -> DQNHyperparameters { let mut p = DQNHyperparameters::conservative(); p.use_branching = true; // Production default — branching kernel is the tested path on H100 p.batch_size = 16; - p.buffer_size = 500; + p.buffer_size = 1024; // MIN_GPU_CAPACITY — GPU PER mandatory, no CPU fallback p.min_replay_size = 32; p.epochs = 3; p.warmup_steps = 0; @@ -36,11 +36,11 @@ pub(super) fn smoke_params() -> DQNHyperparameters { // Set explicit hidden_dim_base to skip GPU capability probe p.hidden_dim_base = Some(32); // Shrink GPU experience collection to fit the small test buffer. - // buffer_size=500, so keep total experiences per collection ≤ buffer_size. + // buffer_size=1024, so keep total experiences per collection ≤ buffer_size. p.gpu_n_episodes = 2; p.gpu_timesteps_per_episode = 50; p.curiosity_weight = 0.0; - // Disable AutoReplaySizer — smoke tests set explicit buffer_size=500. + // Disable AutoReplaySizer — smoke tests set explicit buffer_size=1024. // Without this, H100 (75 GB VRAM) inflates buffer to 10M, causing empty-sample // failures (states length 0) and multi-hour hangs on GPU PER sampling. p.replay_buffer_vram_fraction = 0.0; diff --git a/crates/ml/src/trainers/dqn/trainer/action.rs b/crates/ml/src/trainers/dqn/trainer/action.rs index 2f381196a..3b61511f1 100644 --- a/crates/ml/src/trainers/dqn/trainer/action.rs +++ b/crates/ml/src/trainers/dqn/trainer/action.rs @@ -7,7 +7,6 @@ use tracing::{debug, info}; use super::DQNTrainer; use crate::dqn::action_space::{ExposureLevel, FactoredAction}; use crate::dqn::TradingState; -use crate::dqn::mixed_precision::training_dtype; use crate::dqn::order_router::OrderRouter; use ml_core::fill_simulator::FillResult; @@ -19,7 +18,7 @@ impl DQNTrainer { // Convert state to tensor with tensor core alignment padding let state_vec = state.to_vector(); let raw_dim = state_vec.len(); - let aligned_dim = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_dim, &self.device); + let aligned_dim = (raw_dim + 7) & !7; let padded: Vec = if aligned_dim > raw_dim { let mut v = state_vec.to_vec(); v.resize(aligned_dim, 0.0); @@ -118,7 +117,7 @@ impl DQNTrainer { .map(|s| s.to_vector()) .ok_or_else(|| anyhow::anyhow!("Empty states slice"))?; let raw_state_dim = first_vec.len(); - let aligned_dim = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_state_dim, &self.device); + let aligned_dim = (raw_state_dim + 7) & !7; let pad = aligned_dim - raw_state_dim; // Pre-allocate flat buffer, zero-padding each state to aligned dimension @@ -142,7 +141,7 @@ impl DQNTrainer { // Create batched tensor directly from flat buffer let batch_tensor = Tensor::from_vec(flat_states, (batch_size, aligned_dim), &self.device) .map_err(|e| anyhow::anyhow!("Failed to create batched state tensor: {}", e))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| anyhow::anyhow!("Failed to cast batched state tensor to training dtype: {}", e))?; // FIX: Use get_effective_epsilon() which respects noisy_epsilon_floor. diff --git a/crates/ml/src/trainers/dqn/trainer/constructor.rs b/crates/ml/src/trainers/dqn/trainer/constructor.rs index 2c64eb127..5bfc297b4 100644 --- a/crates/ml/src/trainers/dqn/trainer/constructor.rs +++ b/crates/ml/src/trainers/dqn/trainer/constructor.rs @@ -132,54 +132,55 @@ impl DQNTrainer { // // Guard: skip auto-sizing when buffer_size < 100K (the sizer's own minimum). // A small explicit buffer_size means the caller wants that exact size - // (e.g., smoke tests with buffer_size=500). Without this, H100's 75 GB VRAM - // inflates 500 → 10M, causing empty-sample failures and multi-hour hangs. + // (e.g., smoke tests with buffer_size=1024). Without this, H100's 75 GB VRAM + // inflates small buffers → 10M, causing empty-sample failures and multi-hour hangs. let original_buffer_size = hyperparams.buffer_size; // Save before AutoReplaySizer mutates it - let mut per_max_memory_bytes: usize = 4 * 1024 * 1024 * 1024; // CPU fallback: 4 GB const AUTO_REPLAY_MIN_THRESHOLD: usize = 100_000; + + // Always probe VRAM for PER budget — no hardcoded fallback. + let mut per_max_memory_bytes: usize = if device.is_cuda() { + use ml_core::memory_optimization::detect_gpu_hardware; + match detect_gpu_hardware() { + Ok(hw) => hw.per_max_buffer_bytes(), + Err(e) => { + return Err(anyhow::anyhow!( + "GPU PER requires VRAM probe — detect_gpu_hardware failed: {}", e + )); + } + } + } else { + // CPU-only path (non-CUDA builds): PER budget irrelevant + 0 + }; + + // Dynamic replay buffer sizing: scale replay capacity to available VRAM. + // Only activates when replay_buffer_vram_fraction > 0, GPU detected, and + // buffer_size >= 100K (sizer's own minimum — smaller values are intentional). if hyperparams.replay_buffer_vram_fraction > 0.0 && device.is_cuda() && hyperparams.buffer_size >= AUTO_REPLAY_MIN_THRESHOLD { - use ml_core::memory_optimization::detect_gpu_hardware; - match detect_gpu_hardware() { - Ok(hw) => { - let raw_sd = if hyperparams.mbp10_data_dir.is_some() { 53 } else { 45 }; - let aligned_sd = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_sd, &device); - let replay_cfg = hw.optimal_replay_config( - aligned_sd, - hyperparams.replay_buffer_vram_fraction, - ); - per_max_memory_bytes = replay_cfg.per_max_buffer_bytes; - if replay_cfg.capacity != hyperparams.buffer_size { - info!( - "AutoReplaySizer: replay buffer {} -> {} (VRAM={:.0}MB, fraction={:.0}%, PER budget={:.0}MB)", - hyperparams.buffer_size, - replay_cfg.capacity, - hw.free_memory_mb, - hyperparams.replay_buffer_vram_fraction * 100.0, - per_max_memory_bytes as f64 / (1024.0 * 1024.0), - ); - hyperparams.buffer_size = replay_cfg.capacity; - } - } - Err(e) => { - info!( - "AutoReplaySizer unavailable ({}), using static buffer_size: {}", - e, hyperparams.buffer_size - ); - } - } - } else if device.is_cuda() { - // No auto-sizer, but still compute PER budget from VRAM use ml_core::memory_optimization::detect_gpu_hardware; if let Ok(hw) = detect_gpu_hardware() { - per_max_memory_bytes = hw.per_max_buffer_bytes(); - } else { - // GPU detection failed — keep CPU default (4 GB) + let raw_sd = if hyperparams.mbp10_data_dir.is_some() { 53 } else { 45 }; + let aligned_sd = (raw_sd + 7) & !7; + let replay_cfg = hw.optimal_replay_config( + aligned_sd, + hyperparams.replay_buffer_vram_fraction, + ); + per_max_memory_bytes = replay_cfg.per_max_buffer_bytes; + if replay_cfg.capacity != hyperparams.buffer_size { + info!( + "AutoReplaySizer: replay buffer {} -> {} (VRAM={:.0}MB, fraction={:.0}%, PER budget={:.0}MB)", + hyperparams.buffer_size, + replay_cfg.capacity, + hw.free_memory_mb, + hyperparams.replay_buffer_vram_fraction * 100.0, + per_max_memory_bytes as f64 / (1024.0 * 1024.0), + ); + hyperparams.buffer_size = replay_cfg.capacity; + } } - } else { - // CPU device — keep default 4 GB PER budget } info!( @@ -187,23 +188,6 @@ impl DQNTrainer { if device.is_cuda() { "CUDA GPU" } else { "CPU" }, ); - // Auto-detect mixed precision capability based on GPU architecture - let mixed_precision_detected = if device.is_cuda() { - match crate::memory_optimization::auto_batch_size::detect_gpu_memory() { - Ok((_total, _free, ref name)) => { - let detected = crate::dqn::mixed_precision::detect_from_gpu_name(name); - match &detected { - Some(c) => info!("GPU mixed precision: {:?} enabled (GPU: {})", c.dtype, name), - None => info!("GPU mixed precision: disabled (GPU: {})", name), - } - detected - } - Err(_) => None, - } - } else { - None - }; - // Create DQN configuration // 42-feature architecture: OHLCV, technical, patterns, volume, time, statistical, regime // Portfolio features (3) are populated via PortfolioTracker → 45 total state_dim @@ -215,7 +199,7 @@ impl DQNTrainer { // data pipeline boundaries (GpuPreloadedData and train_batch CPU path). let ofi_enabled = hyperparams.mbp10_data_dir.is_some(); let raw_state_dim = if ofi_enabled { 53 } else { 45 }; - let state_dim = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_state_dim, &device); + let state_dim = (raw_state_dim + 7) & !7; let config = DQNConfig { state_dim, num_actions: 5, // 5 exposure levels (Short100, Short50, Flat, Long50, Long100) @@ -248,7 +232,7 @@ impl DQNTrainer { // PER configuration initial_capital: hyperparams.initial_capital as f64, use_per: hyperparams.use_per, - use_gpu_replay_buffer: hyperparams.use_gpu_replay_buffer, + per_alpha: hyperparams.per_alpha, per_beta_start: hyperparams.per_beta_start, per_beta_max: 1.0, @@ -297,7 +281,6 @@ impl DQNTrainer { minimum_profit_factor: hyperparams.minimum_profit_factor as f32, weight_decay: hyperparams.weight_decay, dropout_rate: if hyperparams.enable_dropout_scheduler { hyperparams.dropout_initial } else { 0.0 }, - mixed_precision: hyperparams.mixed_precision.clone().or(mixed_precision_detected), entropy_coefficient: hyperparams.entropy_coefficient.unwrap_or(0.01), noisy_epsilon_floor: hyperparams.noisy_epsilon_floor.unwrap_or(0.0) as f32, // C2: NoisyNet handles exploration use_count_bonus: hyperparams.count_bonus_coefficient.unwrap_or(0.0) > 0.0, // C3 FIX: enable when coefficient > 0 diff --git a/crates/ml/src/trainers/dqn/trainer/metrics.rs b/crates/ml/src/trainers/dqn/trainer/metrics.rs index 66d630424..55daa8b38 100644 --- a/crates/ml/src/trainers/dqn/trainer/metrics.rs +++ b/crates/ml/src/trainers/dqn/trainer/metrics.rs @@ -4,7 +4,6 @@ use anyhow::Result; use candle_core::Tensor; use super::DQNTrainer; use crate::dqn::TradingState; -use crate::dqn::mixed_precision::training_dtype; use crate::TrainingMetrics; use super::super::config::DQNAgentType; use super::super::statistics::QValueStats; @@ -57,27 +56,14 @@ impl DQNTrainer { // Sample experiences from replay buffer let batch_sample = agent.memory().sample(sample_size)?; - // GPU PER path: use gpu_batch.states directly when available. - // CPU PER fallback: construct tensor from experience Vec. - let batch_tensor = if let Some(ref gpu_batch) = batch_sample.gpu_batch { - gpu_batch.states.to_dtype(training_dtype(agent.device())) - .map_err(|e| crate::MLError::ModelError(format!("GPU Q-stat states dtype cast: {}", e)))? - } else { - let state_dim = agent.get_state_dim(); - let states_flat: Vec = batch_sample.experiences.iter() - .flat_map(|e| { - let mut s = e.state.clone(); - s.resize(state_dim, 0.0); - s - }) - .collect(); - let n = batch_sample.experiences.len(); - candle_core::Tensor::new(states_flat.as_slice(), agent.device()) - .and_then(|t| t.reshape((n, state_dim))) - .and_then(|t| t.to_dtype(training_dtype(agent.device()))) - .map_err(|e| crate::MLError::ModelError(format!("CPU Q-stat states tensor: {}", e)))? - }; - + // GPU PER path: use gpu_batch.states directly (always active in CUDA builds) + let gpu_batch = batch_sample.gpu_batch.as_ref() + .ok_or_else(|| crate::MLError::TrainingError( + "GPU PER must be active — gpu_batch is None".to_owned() + ))?; + let batch_tensor = gpu_batch.states.to_dtype(candle_core::DType::BF16) + .map_err(|e| crate::MLError::ModelError(format!("GPU Q-stat states dtype cast: {}", e)))?; + // Forward pass to get Q-values [batch_size, num_actions] let q_values = agent.forward(&batch_tensor)?; @@ -302,7 +288,7 @@ impl DQNTrainer { { if let Some(ref gpu) = batch_sample.gpu_batch { batch_tensor_opt = gpu.states - .to_dtype(training_dtype(agent.device())) + .to_dtype(candle_core::DType::BF16) .ok(); } } @@ -321,7 +307,7 @@ impl DQNTrainer { let t = match Tensor::from_vec(batched_states, (sample_size, state_dim), agent.device()) { Ok(t) => t, Err(_) => return None, }; - match t.to_dtype(training_dtype(agent.device())) { + match t.to_dtype(candle_core::DType::BF16) { Ok(t) => t, Err(_) => return None, } @@ -429,7 +415,7 @@ fn compute_q_diagnostics_gpu( let agent = self.agent.read().await; let state_vec = state.to_vector(); let raw_dim = state_vec.len(); - let aligned = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_dim, &self.device); + let aligned = (raw_dim + 7) & !7; let padded: Vec = if aligned > raw_dim { let mut v = state_vec.to_vec(); v.resize(aligned, 0.0); @@ -603,11 +589,11 @@ fn compute_q_diagnostics_gpu( .map_err(|e| anyhow::anyhow!("GPU val pad zeros: {e}"))?; Tensor::cat(&[&state_gpu, &pad], 1) .map_err(|e| anyhow::anyhow!("GPU val state pad: {e}"))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| anyhow::anyhow!("GPU val state dtype: {e}"))? } else { state_gpu - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| anyhow::anyhow!("GPU val state dtype: {e}"))? }; @@ -761,26 +747,11 @@ fn compute_q_diagnostics_gpu( .sample(sample_size) .map_err(|e| anyhow::anyhow!("Failed to sample experiences: {}", e))?; - // GPU PER path: use gpu_batch.states directly when available. - // CPU PER fallback: construct tensor from experience Vec. - let batch_tensor = if let Some(ref gpu_batch) = batch_sample.gpu_batch { - gpu_batch.states.to_dtype(training_dtype(agent.device())) - .map_err(|e| anyhow::anyhow!("GPU Q-est states dtype cast: {}", e))? - } else { - let state_dim = agent.get_state_dim(); - let states_flat: Vec = batch_sample.experiences.iter() - .flat_map(|e| { - let mut s = e.state.clone(); - s.resize(state_dim, 0.0); - s - }) - .collect(); - let n = batch_sample.experiences.len(); - candle_core::Tensor::new(states_flat.as_slice(), agent.device()) - .and_then(|t| t.reshape((n, state_dim))) - .and_then(|t| t.to_dtype(training_dtype(agent.device()))) - .map_err(|e| anyhow::anyhow!("CPU Q-est states tensor: {}", e))? - }; + // GPU PER path: use gpu_batch.states directly (always active in CUDA builds) + let gpu_batch = batch_sample.gpu_batch.as_ref() + .ok_or_else(|| anyhow::anyhow!("GPU PER must be active — gpu_batch is None"))?; + let batch_tensor = gpu_batch.states.to_dtype(candle_core::DType::BF16) + .map_err(|e| anyhow::anyhow!("GPU Q-est states dtype cast: {}", e))?; // WAVE 23 P0 Fix: Check for Q-value divergence (early stopping) // This calls log_q_values() which returns Err if divergence detected for consecutive checks diff --git a/crates/ml/src/trainers/dqn/trainer/state.rs b/crates/ml/src/trainers/dqn/trainer/state.rs index 259d84395..0e0a3d5fb 100644 --- a/crates/ml/src/trainers/dqn/trainer/state.rs +++ b/crates/ml/src/trainers/dqn/trainer/state.rs @@ -131,7 +131,7 @@ impl DQNTrainer { // Convert TradingState to flat vector, pad for tensor core alignment let state_vec = trading_state.to_vector(); let raw_dim = state_vec.len(); - let aligned = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_dim, &self.device); + let aligned = (raw_dim + 7) & !7; let padded: Vec = if aligned > raw_dim { let mut v = state_vec.to_vec(); v.resize(aligned, 0.0); diff --git a/crates/ml/src/trainers/dqn/trainer/tests.rs b/crates/ml/src/trainers/dqn/trainer/tests.rs index 831b50e47..e09e6abdf 100644 --- a/crates/ml/src/trainers/dqn/trainer/tests.rs +++ b/crates/ml/src/trainers/dqn/trainer/tests.rs @@ -55,10 +55,7 @@ fn create_test_trainer_with(params: DQNHyperparameters) -> Result { /// Pad a TradingState's regime_features so that `state.dimension()` matches the /// trainer's aligned state_dim (e.g. 45→48 on CUDA due to tensor core alignment). fn pad_state_to_aligned(state: &mut TradingState, trainer: &DQNTrainer) { - let aligned_dim = crate::dqn::mixed_precision::align_dim_for_tensor_cores( - state.dimension(), - &trainer.device, - ); + let aligned_dim = (state.dimension() + 7) & !7; let pad = aligned_dim.saturating_sub(state.dimension()); if pad > 0 { state.regime_features.extend(vec![0.0_f32; pad]); @@ -437,7 +434,7 @@ async fn test_train_with_empty_data_completes_gracefully() { params.epochs = 5; // Short run — just checking it doesn't panic params.early_stopping_enabled = false; params.gradient_collapse_patience = 1000; - params.buffer_size = 1000; + params.buffer_size = 1024; // MIN_GPU_CAPACITY — GPU PER mandatory let device = candle_core::Device::new_cuda(0).expect("CUDA device required"); let mut trainer = DQNTrainer::new_with_device(params, device).unwrap(); let empty_data: Vec<(FeatureVector, Vec)> = vec![]; diff --git a/crates/ml/src/trainers/dqn/trainer/train_step.rs b/crates/ml/src/trainers/dqn/trainer/train_step.rs index 00103e7e5..e92ee3b65 100644 --- a/crates/ml/src/trainers/dqn/trainer/train_step.rs +++ b/crates/ml/src/trainers/dqn/trainer/train_step.rs @@ -4,7 +4,6 @@ use anyhow::Result; use candle_core::{IndexOp, Tensor}; use tracing::{debug, info, warn}; -use crate::dqn::mixed_precision::training_dtype; use super::DQNTrainer; impl DQNTrainer { @@ -183,7 +182,7 @@ impl DQNTrainer { if let Some(ref gpu) = batch_sample.gpu_batch { batch_tensor_opt = Some( gpu.states - .to_dtype(training_dtype(agent.device())) + .to_dtype(candle_core::DType::BF16) .map_err(|e| anyhow::anyhow!("Q-est dtype: {e}"))?, ); } @@ -199,7 +198,7 @@ impl DQNTrainer { &self.device, ) .map_err(|e| anyhow::anyhow!("Q-est tensor: {e}"))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| anyhow::anyhow!("Q-est dtype: {e}"))?; batch_tensor_opt = Some(tensor); } @@ -512,7 +511,7 @@ impl DQNTrainer { if let Some(ref gpu) = batch_sample.gpu_batch { batch_tensor_opt = Some( gpu.states - .to_dtype(training_dtype(agent.device())) + .to_dtype(candle_core::DType::BF16) .map_err(|e| anyhow::anyhow!("Q-est dtype: {e}"))?, ); } @@ -528,7 +527,7 @@ impl DQNTrainer { &self.device, ) .map_err(|e| anyhow::anyhow!("Q-est tensor: {e}"))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| anyhow::anyhow!("Q-est dtype: {e}"))?; batch_tensor_opt = Some(tensor); } @@ -603,10 +602,6 @@ impl DQNTrainer { /// Lazy-init fused CUDA training context (Standard DQN only). /// Recreate if batch_size changed (OOM recovery). - /// - /// Requires GPU PER — CUDA Graph capture cannot include CPU→GPU transfers - /// from CPU PER sampling. When GPU PER is unavailable (e.g. small test - /// buffers), we skip fused init and fall through to the Candle training path. pub(crate) async fn ensure_fused_ctx(&mut self) { let needs_init = match &self.fused_ctx { None => self.device.is_cuda(), @@ -615,17 +610,6 @@ impl DQNTrainer { if !needs_init { return; } - // Fused CUDA Graphs require GPU-resident PER for stream-capture-compatible - // sampling. Skip fused init when agent fell back to CPU PER. - #[cfg(feature = "cuda")] - { - let agent_check = self.agent.read().await; - if !agent_check.has_gpu_per() { - debug!("Fused CUDA training skipped: GPU PER not available (CPU PER fallback active)"); - return; - } - drop(agent_check); - } if self.fused_ctx.is_some() { info!("Fused CUDA context: batch_size changed, recreating"); self.fused_ctx = None; diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index e9e9cb2f3..08a2c429b 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -349,7 +349,7 @@ impl DQNTrainer { // Set tensor core alignment so build_*_states pads output let ofi_enabled = self.ofi_features.is_some(); let raw_dim = if ofi_enabled { 53 } else { 45 }; - let aligned_dim = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_dim, &self.device); + let aligned_dim = (raw_dim + 7) & !7; gpu_data.set_aligned_state_dim(aligned_dim); self.gpu_data = Some(gpu_data); } @@ -661,7 +661,7 @@ impl DQNTrainer { use crate::cuda_pipeline::gpu_experience_collector::ExperienceCollectorConfig; let raw_sd = if self.hyperparams.mbp10_data_dir.is_some() { 53 } else { 45 }; - let aligned_sd = crate::dqn::mixed_precision::align_dim_for_tensor_cores(raw_sd, &self.device); + let aligned_sd = (raw_sd + 7) & !7; // Cache n_episodes on first epoch let n_episodes = if let Some(cached) = self.cached_n_episodes { diff --git a/crates/ml/src/trainers/online_learning.rs b/crates/ml/src/trainers/online_learning.rs index cfe7513e5..ebce05bf9 100644 --- a/crates/ml/src/trainers/online_learning.rs +++ b/crates/ml/src/trainers/online_learning.rs @@ -600,7 +600,7 @@ mod tests { fn tiny_var_map() -> Result { let dev = cuda_device(); let var_map = VarMap::new(); - let vb = candle_nn::VarBuilder::from_varmap(&var_map, crate::dqn::mixed_precision::training_dtype(&dev), &dev); + let vb = candle_nn::VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &dev); let _linear = candle_nn::linear(2, 2, vb.pp("layer")) .map_err(|e| MLError::ModelError(format!("tiny_var_map linear: {e}")))?; Ok(var_map) diff --git a/crates/ml/src/trainers/ppo.rs b/crates/ml/src/trainers/ppo.rs index 280b63b93..e12f1667e 100644 --- a/crates/ml/src/trainers/ppo.rs +++ b/crates/ml/src/trainers/ppo.rs @@ -191,7 +191,6 @@ impl From for PPOConfig { lstm_sequence_length: 32, accumulation_steps: params.accumulation_steps.max(1), clip_epsilon_high: None, - mixed_precision: None, // Auto-detected at trainer initialization use_symlog: true, use_adaptive_entropy: true, use_percentile_scaling: true, @@ -366,18 +365,6 @@ impl PpoTrainer { config.gae_config.gamma = hyperparams.gamma as f32; config.gae_config.lambda = hyperparams.gae_lambda; - // Auto-detect mixed precision capability based on GPU architecture - if device.is_cuda() { - if let Ok((_total, _free, ref name)) = crate::memory_optimization::auto_batch_size::detect_gpu_memory() { - let detected = crate::dqn::mixed_precision::detect_from_gpu_name(name); - match &detected { - Some(c) => info!("PPO mixed precision: {:?} enabled (GPU: {})", c.dtype, name), - None => info!("PPO mixed precision: disabled (GPU: {})", name), - } - config.mixed_precision = detected; - } - } - // Create PPO model with specified device (GPU or CPU) let model = PPO::with_device(config, device.clone())?; diff --git a/crates/ml/src/trainers/tft/config.rs b/crates/ml/src/trainers/tft/config.rs index df7cb7ea2..bde9e19b2 100644 --- a/crates/ml/src/trainers/tft/config.rs +++ b/crates/ml/src/trainers/tft/config.rs @@ -140,7 +140,6 @@ impl TFTTrainerConfig { dropout_rate: self.dropout_rate, l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: true, // Auto-detected per GPU capabilities memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000, diff --git a/crates/ml/src/trainers/tlob.rs b/crates/ml/src/trainers/tlob.rs index 49f350005..6101c2364 100644 --- a/crates/ml/src/trainers/tlob.rs +++ b/crates/ml/src/trainers/tlob.rs @@ -33,7 +33,6 @@ use tracing::{info, instrument, warn}; use candle_nn::ParamsAdamW; -use crate::dqn::mixed_precision::training_dtype; use crate::tlob::features::TLOB_FEATURE_COUNT; use crate::tlob::transformer::TLOBTransformer; @@ -209,7 +208,7 @@ impl TLOBTrainer { // Initialize variable map and var builder let var_map = Arc::new(VarMap::new()); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(&device), &device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, &device); // Create TLOB transformer model (ONNX-based inference / fallback) let model = Self::create_trainable_model(&hyperparams, vb.clone(), &device)?; diff --git a/crates/ml/src/training/orchestrator.rs b/crates/ml/src/training/orchestrator.rs index 52f7871e0..e97839176 100644 --- a/crates/ml/src/training/orchestrator.rs +++ b/crates/ml/src/training/orchestrator.rs @@ -30,8 +30,6 @@ pub struct OrchestratorConfig { pub lr_schedule: LRSchedule, /// Enable gradient accumulation pub gradient_accumulation_steps: usize, - /// Enable mixed precision training (if supported) - pub mixed_precision: bool, /// Maximum gradient norm for clipping pub max_grad_norm: Option, } @@ -46,7 +44,6 @@ impl Default for OrchestratorConfig { early_stopping_patience: Some(10), lr_schedule: LRSchedule::Constant, gradient_accumulation_steps: 1, - mixed_precision: true, // Auto-detected per GPU capabilities max_grad_norm: Some(1.0), } } diff --git a/crates/ml/src/training_pipeline.rs b/crates/ml/src/training_pipeline.rs index 103c450e0..4467d20e5 100644 --- a/crates/ml/src/training_pipeline.rs +++ b/crates/ml/src/training_pipeline.rs @@ -173,8 +173,6 @@ pub struct PerformanceConfig { pub device_preference: String, /// Maximum memory usage (bytes) pub max_memory_bytes: usize, - /// Enable mixed precision training - pub mixed_precision: bool, /// Number of data loader workers pub num_workers: usize, /// Enable gradient accumulation @@ -788,7 +786,6 @@ impl Default for ProductionTrainingConfig { performance_config: PerformanceConfig { device_preference: "cpu".to_owned(), max_memory_bytes: 8 * 1024 * 1024 * 1024, // 8GB - mixed_precision: true, // Auto-detected per GPU capabilities num_workers: 4, gradient_accumulation_steps: 1, }, diff --git a/crates/ml/src/validation/ppo_adapter.rs b/crates/ml/src/validation/ppo_adapter.rs index b590649dd..8abd5c9f2 100644 --- a/crates/ml/src/validation/ppo_adapter.rs +++ b/crates/ml/src/validation/ppo_adapter.rs @@ -11,7 +11,7 @@ use std::cell::RefCell; use candle_core::{DType, Device, Tensor}; -use ml_core::mixed_precision::training_dtype; + use rand::Rng; use crate::common::action::FactoredAction; @@ -99,7 +99,7 @@ impl ValidatableStrategy for PpoStrategy { &self.device, ) .map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast state to training dtype: {}", e)))?; let ppo = self.ppo.borrow(); @@ -152,7 +152,7 @@ impl ValidatableStrategy for PpoStrategy { &self.device, ) .map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast state to training dtype: {}", e)))?; let ppo = self.ppo.borrow(); @@ -275,7 +275,7 @@ impl ValidatableStrategy for PpoLstmStrategy { &self.device, ) .map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast state to training dtype: {}", e)))?; // Get hidden states @@ -367,7 +367,7 @@ impl ValidatableStrategy for PpoLstmStrategy { &self.device, ) .map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))? - .to_dtype(training_dtype(&self.device)) + .to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(format!("Failed to cast state to training dtype: {}", e)))?; let (policy_h, policy_c) = hsm.get_policy_state(); diff --git a/crates/ml/src/xlstm/trainable.rs b/crates/ml/src/xlstm/trainable.rs index 2d0130eee..9489f0ed2 100644 --- a/crates/ml/src/xlstm/trainable.rs +++ b/crates/ml/src/xlstm/trainable.rs @@ -9,7 +9,7 @@ use std::collections::HashMap; use super::config::XLSTMConfig; use super::network::XLSTMNetwork; -use crate::dqn::mixed_precision::training_dtype; + use crate::training::unified_trainer::{checkpoint, CheckpointMetadata, TrainingMetrics, UnifiedTrainable}; use crate::MLError; @@ -43,7 +43,7 @@ impl XLSTMTrainableAdapter { /// Create a new xLSTM trainable adapter. pub fn new(config: XLSTMConfig, device: &Device) -> Result { let var_map = VarMap::new(); - let vb = VarBuilder::from_varmap(&var_map, training_dtype(device), device); + let vb = VarBuilder::from_varmap(&var_map, candle_core::DType::BF16, device); let network = XLSTMNetwork::new(&config, vb)?; @@ -88,7 +88,7 @@ impl UnifiedTrainable for XLSTMTrainableAdapter { } fn forward(&mut self, input: &Tensor) -> Result { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = input.to_dtype(candle_core::DType::BF16) .map_err(|e| MLError::ModelError(e.to_string()))?; self.network.forward(&input) } diff --git a/crates/ml/tests/dqn_accumulation_convergence_test.rs b/crates/ml/tests/dqn_accumulation_convergence_test.rs index 2ea8b34d1..4b3e3e35e 100644 --- a/crates/ml/tests/dqn_accumulation_convergence_test.rs +++ b/crates/ml/tests/dqn_accumulation_convergence_test.rs @@ -109,6 +109,7 @@ async fn test_accumulation_convergence_similar_to_direct() -> Result<()> { // --- Run 1: Accumulated training (batch=16, accum=4) --- let mut hp_accum = DQNHyperparameters::conservative(); + hp_accum.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hp_accum.epochs = epochs; hp_accum.batch_size = 16; hp_accum.gradient_accumulation_steps = 4; @@ -129,6 +130,7 @@ async fn test_accumulation_convergence_similar_to_direct() -> Result<()> { // --- Run 2: Direct training (batch=64, accum=1) --- let mut hp_direct = DQNHyperparameters::conservative(); + hp_direct.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hp_direct.epochs = epochs; hp_direct.batch_size = 64; hp_direct.gradient_accumulation_steps = 1; diff --git a/crates/ml/tests/dqn_gradient_accumulation_test.rs b/crates/ml/tests/dqn_gradient_accumulation_test.rs index 3d2b6bd41..20ad5506c 100644 --- a/crates/ml/tests/dqn_gradient_accumulation_test.rs +++ b/crates/ml/tests/dqn_gradient_accumulation_test.rs @@ -111,6 +111,7 @@ async fn test_accumulation_single_optimizer_step() -> Result<()> { let checkpoint_dir = tempfile::tempdir()?; let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 5; hyperparams.batch_size = 32; hyperparams.learning_rate = 0.0001; diff --git a/crates/ml/tests/dqn_inference_test.rs b/crates/ml/tests/dqn_inference_test.rs index dd8ffd495..2e044270c 100644 --- a/crates/ml/tests/dqn_inference_test.rs +++ b/crates/ml/tests/dqn_inference_test.rs @@ -172,6 +172,7 @@ async fn test_checkpoint_to_inference() -> Result<()> { // Phase 1: Train for 5 epochs to produce a valid checkpoint // ===================================================================== let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 5; hyperparams.batch_size = 32; hyperparams.learning_rate = 0.0001; diff --git a/crates/ml/tests/dqn_long_training_test.rs b/crates/ml/tests/dqn_long_training_test.rs index 3903d98ec..a84568f4a 100644 --- a/crates/ml/tests/dqn_long_training_test.rs +++ b/crates/ml/tests/dqn_long_training_test.rs @@ -122,6 +122,7 @@ async fn test_dqn_50_epoch_convergence() -> Result<()> { // --- Configure hyperparameters --- let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 50; hyperparams.batch_size = 64; hyperparams.learning_rate = 0.0001; diff --git a/crates/ml/tests/dqn_training_pipeline_test.rs b/crates/ml/tests/dqn_training_pipeline_test.rs index 8af65bea0..11cccd621 100644 --- a/crates/ml/tests/dqn_training_pipeline_test.rs +++ b/crates/ml/tests/dqn_training_pipeline_test.rs @@ -161,6 +161,7 @@ async fn test_dqn_trains_on_es_fut() -> Result<()> { // Configure hyperparameters for fast test (5 epochs) let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 3; // CI smoke: validate pipeline, not convergence hyperparams.batch_size = 64; hyperparams.learning_rate = 0.001; @@ -287,6 +288,7 @@ async fn test_dqn_loss_decreases() -> Result<()> { // Train for 3 epochs — CI validates gradient flow, not full convergence let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 3; hyperparams.batch_size = 64; hyperparams.learning_rate = 0.001; @@ -361,6 +363,7 @@ async fn test_dqn_checkpoint_save_load() -> Result<()> { // Train for 2 epochs and save checkpoint let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 2; hyperparams.batch_size = 64; hyperparams.checkpoint_frequency = 2; @@ -430,6 +433,7 @@ async fn test_dqn_q_value_predictions() -> Result<()> { // Train minimal model let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 2; hyperparams.batch_size = 32; @@ -487,6 +491,7 @@ async fn test_dqn_epsilon_greedy() -> Result<()> { // Configure with high epsilon decay let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 2; hyperparams.epsilon_start = 1.0; hyperparams.epsilon_end = 0.01; @@ -560,6 +565,7 @@ async fn test_dqn_full_production_training() -> Result<()> { // Production hyperparameters let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 50; hyperparams.batch_size = 128; hyperparams.learning_rate = 0.0001; diff --git a/crates/ml/tests/dqn_training_smoke_test.rs b/crates/ml/tests/dqn_training_smoke_test.rs index ff1ebef3c..68b5a5987 100644 --- a/crates/ml/tests/dqn_training_smoke_test.rs +++ b/crates/ml/tests/dqn_training_smoke_test.rs @@ -129,6 +129,7 @@ async fn test_dqn_training_smoke() -> Result<()> { let checkpoint_dir = tempfile::tempdir()?; let mut hyperparams = DQNHyperparameters::conservative(); + hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism hyperparams.epochs = 3; hyperparams.batch_size = 256; // H100 80GB: saturate tensor cores hyperparams.learning_rate = 0.001; // Must outpace soft target updates (tau=0.001) under noisy nets diff --git a/crates/ml/tests/gpu_per_integration_test.rs b/crates/ml/tests/gpu_per_integration_test.rs index 22d2d6765..5403ae3e4 100644 --- a/crates/ml/tests/gpu_per_integration_test.rs +++ b/crates/ml/tests/gpu_per_integration_test.rs @@ -325,13 +325,13 @@ fn test_gpu_per_ring_buffer_overwrites_correctly() { #[test] fn test_gpu_per_oom_rejects_absurd_capacity() { - // Absurd allocation that exceeds both GPU and CPU memory limits. + // Absurd allocation that exceeds GPU memory limits. // Must return Err — not panic or OOM-kill the system. let device = Device::new_cuda(0).expect("CUDA required"); - // 500M capacity × 4 state_dim: GPU pre-flight rejects (~24 GB), - // CPU fallback pre-flight also rejects (~34 GB estimated). - let result = ReplayBufferType::try_gpu_prioritized_with_fallback( + // 500M capacity × 4 state_dim: GPU pre-flight rejects (~24 GB). + // No CPU fallback — GPU PER is mandatory. + let result = ReplayBufferType::try_gpu_with_halving( 500_000_000, // 500M capacity 4, // 4 state_dim 0.6, @@ -344,6 +344,6 @@ fn test_gpu_per_oom_rejects_absurd_capacity() { assert!( result.is_err(), - "500M capacity should be rejected by both GPU and CPU pre-flight checks" + "500M capacity should be rejected by GPU pre-flight checks" ); } diff --git a/crates/ml/tests/ppo_45_action_network_tests.rs b/crates/ml/tests/ppo_45_action_network_tests.rs index 7b90a8c22..587b1ca33 100644 --- a/crates/ml/tests/ppo_45_action_network_tests.rs +++ b/crates/ml/tests/ppo_45_action_network_tests.rs @@ -368,7 +368,6 @@ fn test_hyperopt_adapter_default_45_actions() -> Result<()> { lstm_sequence_length: 32, accumulation_steps: 1, clip_epsilon_high: None, - mixed_precision: None, use_symlog: true, use_adaptive_entropy: true, use_percentile_scaling: true, diff --git a/crates/ml/tests/ppo_recurrent_integration_tests.rs b/crates/ml/tests/ppo_recurrent_integration_tests.rs index 337e615f1..b9249fd7d 100644 --- a/crates/ml/tests/ppo_recurrent_integration_tests.rs +++ b/crates/ml/tests/ppo_recurrent_integration_tests.rs @@ -135,7 +135,6 @@ fn test_recurrent_ppo_single_episode() { lstm_sequence_length: 32, accumulation_steps: 1, clip_epsilon_high: None, - mixed_precision: None, use_symlog: true, use_adaptive_entropy: true, use_percentile_scaling: true, diff --git a/crates/ml/tests/tft_inference_latency_benchmark.rs b/crates/ml/tests/tft_inference_latency_benchmark.rs index a091c0a4c..3c0773179 100644 --- a/crates/ml/tests/tft_inference_latency_benchmark.rs +++ b/crates/ml/tests/tft_inference_latency_benchmark.rs @@ -154,7 +154,6 @@ fn test_tft_inference_latency_p95_target() -> Result<(), MLError> { dropout_rate: 0.0, // Inference mode: No dropout l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: false, // Test FP32 baseline first memory_efficient: true, max_inference_latency_us: 5000, target_throughput_pps: 100_000, diff --git a/crates/ml/tests/tft_real_dbn_data_test.rs b/crates/ml/tests/tft_real_dbn_data_test.rs index 4c4705bf1..2aa93dcec 100644 --- a/crates/ml/tests/tft_real_dbn_data_test.rs +++ b/crates/ml/tests/tft_real_dbn_data_test.rs @@ -495,7 +495,6 @@ fn create_test_tft_config() -> TFTConfig { dropout_rate: 0.1, l2_regularization: 0.0001, use_flash_attention: false, // Disable for compatibility - mixed_precision: false, // F32 for stability memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000, diff --git a/crates/ml/tests/tft_test.rs b/crates/ml/tests/tft_test.rs index ee3ed0f57..90ef36df9 100644 --- a/crates/ml/tests/tft_test.rs +++ b/crates/ml/tests/tft_test.rs @@ -116,7 +116,6 @@ fn test_tft_config_custom() -> Result<()> { dropout_rate: 0.1, l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: true, memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000, @@ -240,7 +239,6 @@ fn test_tft_config_validation() -> Result<()> { dropout_rate: 0.1, l2_regularization: 0.0001, use_flash_attention: false, - mixed_precision: false, memory_efficient: true, max_inference_latency_us: 100, target_throughput_pps: 50_000, @@ -390,7 +388,6 @@ async fn test_tft_config_validation_real_data() -> Result<()> { dropout_rate: 0.1, l2_regularization: 0.0001, use_flash_attention: false, - mixed_precision: false, memory_efficient: true, max_inference_latency_us: 100, target_throughput_pps: 50_000, diff --git a/docs/superpowers/plans/2026-03-15-fused-cuda-training.md b/docs/superpowers/plans/2026-03-15-fused-cuda-training.md new file mode 100644 index 000000000..c8188e557 --- /dev/null +++ b/docs/superpowers/plans/2026-03-15-fused-cuda-training.md @@ -0,0 +1,259 @@ +# Fused CUDA Training Kernel — Implementation Plan + +> **For agentic workers:** REQUIRED: Use superpowers:subagent-driven-development (if subagents available) or superpowers:executing-plans to implement this plan. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Replace 2,100+ Candle kernel dispatches per DQN training batch with 3 fused CUDA kernels captured in a CUDA Graph, achieving ~15-20x epoch speedup on H100. + +**Architecture:** Fused forward+loss, backward, and Adam kernels bypass Candle entirely. Weight pointers extracted from existing DuelingWeightSet/BranchingWeightSet. CUDA Graph captures the fixed-shape kernel sequence for zero-overhead replay. + +**Tech Stack:** CUDA (NVRTC), cudarc 0.17.3, Candle (weight storage only), existing common_device_functions.cuh infrastructure. + +**Spec:** `docs/superpowers/specs/2026-03-15-fused-cuda-training-design.md` + +--- + +## Chunk 1: Forward + Loss Kernel + +### Task 1: Write the CUDA forward+loss kernel + +**Files:** +- Create: `crates/ml/src/cuda_pipeline/dqn_training_kernel.cu` + +**Dependencies:** Uses macros/functions from `common_device_functions.cuh` (TILE_LAYER_WARP_CLEAN, q_forward_dueling_warp_shmem pattern, cooperative_load_tile, warp_matvec_leaky_relu_shmem). + +- [ ] **Step 1: Write kernel entry point and data loading** + +Kernel loads batch data (states, next_states, actions, rewards, dones, IS weights) into registers. +One block (32 threads = 1 warp) per sample. Reuses existing warp-cooperative pattern. + +- [ ] **Step 2: Implement 3 forward passes** + +Extend existing `q_forward_dueling_warp_shmem` pattern: +- Online forward on states (save activations for backward: h_s1, h_s2, h_v1, h_bd) +- Target forward on next_states (no saves) +- Online forward on next_states for Double DQN action selection (no saves) + +For distributional mode: output is [n_d × num_atoms] per branch, apply log_softmax per action. + +- [ ] **Step 3: Implement C51 distributional loss** + +Per-branch cross-entropy: +1. Decompose factored action: exp_idx = action / 9, ord_idx = (action % 9) / 3, urg_idx = action % 3 +2. Gather current log-probs for taken action: index into log_softmax output +3. Select best next action from online forward on next_states (argmax) +4. Gather target probs for best next action +5. Bellman projection: T_z = r + γ×z×(1-done), clip to [v_min, v_max], linear interpolation scatter +6. Cross-entropy: -Σ projected × current_log_probs +7. Average over 3 branches, multiply by IS weight + +- [ ] **Step 4: Write outputs** + +Write per-sample loss, td_errors to global memory. Lane 0 atomicAdd to total_loss. +Save activations to pre-allocated buffers for backward kernel. + +### Task 2: Write the Rust host code for forward+loss kernel + +**Files:** +- Create: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` +- Modify: `crates/ml/src/cuda_pipeline/mod.rs` (add module) + +- [ ] **Step 1: Define GpuDqnTrainer struct with pre-allocated buffers** + +All buffers allocated once at construction (fixed shapes for CUDA Graph compatibility): +- Batch input buffers (states, next_states, actions, rewards, dones, is_weights) +- Activation save buffers (h_s1, h_s2, h_v1, h_bd, logits, target_probs) +- Output buffers (per_sample_loss, td_errors, total_loss) + +- [ ] **Step 2: Implement NVRTC compilation with #define injection** + +Follow existing pattern from `compile_forward_kernel()` in gpu_backtest_evaluator.rs: +- dim_overrides → common_device_functions.cuh → dqn_training_kernel.cu +- Inject: STATE_DIM, SHARED_H1/H2, VALUE_H, ADV_H, NUM_ATOMS, V_MIN, V_MAX, branch sizes, BATCH_SIZE + +- [ ] **Step 3: Implement kernel launch** + +Extract weight pointers from DuelingWeightSet + BranchingWeightSet (online + target). +Launch with grid=(batch_size, 1, 1), block=(32, 1, 1), shared_mem=tile_size. + +- [ ] **Step 4: Test forward+loss numerical correctness** + +Run same batch through Candle path and fused kernel, compare: +- Q-values within 1e-5 relative error +- Per-sample loss within 1e-4 +- td_errors within 1e-4 + +--- + +## Chunk 2: Backward Kernel + +### Task 3: Write the CUDA backward kernel + +**Files:** +- Modify: `crates/ml/src/cuda_pipeline/dqn_training_kernel.cu` + +- [ ] **Step 1: Implement gradient through C51 cross-entropy + log_softmax** + +For each branch d, for taken action a_d: +- ∂CE/∂log_p = -projected_target (shape: [num_atoms]) +- Through log_softmax: ∂L/∂z = ∂L/∂log_p - softmax(z) × Σ(∂L/∂log_p) + +- [ ] **Step 2: Implement backward through linear layers** + +For each linear layer (reverse order): +- ∂L/∂W = outer_product(∂L/∂y, x) → atomicAdd to gradient buffer +- ∂L/∂b = ∂L/∂y → atomicAdd +- ∂L/∂x = matmul(W^T, ∂L/∂y) +- Through LeakyReLU: mask by sign of pre-activation + +Accumulate ∂L/∂h_s2 from value head + all 3 branch heads. + +- [ ] **Step 3: Wire zero-init of gradient buffers** + +Before backward kernel: memset gradient buffers to 0. +Use stream.memset_zeros() (captured in CUDA Graph). + +- [ ] **Step 4: Test backward numerical correctness** + +Compare gradients against Candle backward() within 1e-3 tolerance (atomicAdd accumulation noise). + +### Task 4: Extend Rust host for backward kernel + +**Files:** +- Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` + +- [ ] **Step 1: Add gradient buffer allocation** + +Flattened gradient buffer: [TOTAL_PARAMS] floats. Layout matches parameter flattening order. +Zero-init between batches. + +- [ ] **Step 2: Implement backward kernel launch** + +Pass saved activation buffers, weight pointers (for W^T), gradient output buffers. +Grid/block same as forward kernel. + +--- + +## Chunk 3: Adam Optimizer Kernel + +### Task 5: Write the Adam optimizer kernel + +**Files:** +- Modify: `crates/ml/src/cuda_pipeline/dqn_training_kernel.cu` + +- [ ] **Step 1: Implement gradient norm + clipping** + +Two-pass approach: +- Pass 1: Each thread accumulates grad² for its elements, block reduction, atomicAdd to global norm +- __threadfence() + last-block detection +- Pass 2: scale = min(max_norm / (norm + eps), 1.0), each thread scales its grads + +- [ ] **Step 2: Implement Adam update** + +Per-element (trivially parallel, 256 threads per block, ceil(TOTAL_PARAMS/256) blocks): +- m[i] = β1×m[i] + (1-β1)×g[i] +- v[i] = β2×v[i] + (1-β2)×g[i]² +- m_hat = m[i] / (1-β1^t), v_hat = v[i] / (1-β2^t) +- param[i] -= lr × m_hat / (√v_hat + ε) + wd × param[i] + +- [ ] **Step 3: Output grad_norm** + +Write pre-clip gradient L2 norm to output buffer (for monitoring). + +### Task 6: Extend Rust host for Adam kernel + +**Files:** +- Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` + +- [ ] **Step 1: Add Adam state buffers (m, v)** + +Allocated once, persisted across batches. Same size as gradient buffer. +Initialize to zeros. + +- [ ] **Step 2: Implement parameter flattening/unflattening** + +Map Candle VarMap tensors ↔ flat CUDA buffer for optimizer. +After Adam update, sync back to VarMap tensors for target network EMA. + +- [ ] **Step 3: Test Adam correctness** + +Compare parameter updates against Candle Adam within 1e-4. + +--- + +## Chunk 4: CUDA Graph Capture + +### Task 7: Implement CUDA Graph capture and replay + +**Files:** +- Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` + +- [ ] **Step 1: Implement graph capture** + +On first batch: +```rust +stream.begin_capture(CU_STREAM_CAPTURE_MODE_THREAD_LOCAL)?; +// Launch: noise_gen (if noisy), zero_grad, forward_loss, backward, adam +stream.end_capture(CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH)?; +``` + +- [ ] **Step 2: Implement graph replay loop** + +Subsequent batches: +1. Write new batch data into pre-allocated input buffers (outside graph) +2. Generate new NoisyNet noise (outside graph) +3. graph.launch() — replays entire training step + +- [ ] **Step 3: Implement graph invalidation** + +`invalidate_training_graph()` — called when weights need manual update (target network sync). +Rebuild graph on next batch. + +- [ ] **Step 4: Smoke test: capture + 100 replays** + +Verify loss decreases, no crashes, no memory leaks. + +--- + +## Chunk 5: Integration + +### Task 8: Wire into DQN train_step + +**Files:** +- Modify: `crates/ml-dqn/src/dqn.rs` (train_step method) +- Modify: `crates/ml/src/cuda_pipeline/mod.rs` (register module) + +- [ ] **Step 1: Add gpu_trainer field to DQNAgent** + +Lazily initialized on first CUDA train_step. Requires: network dims, batch_size, config. + +- [ ] **Step 2: Modify train_step to use fused path** + +```rust +#[cfg(feature = "cuda")] +if self.gpu_trainer.is_some() { + // Fused CUDA path — bypasses Candle entirely + return self.train_step_fused(batch); +} +// Fallback: existing Candle path (tests, CPU builds) +``` + +- [ ] **Step 3: Implement weight sync** + +After target network EMA update (Candle VarMap), sync changed weights to CUDA buffers. +Invalidate CUDA Graph (target weights changed). + +- [ ] **Step 4: End-to-end integration test** + +Full training run (1 epoch, 100 batches) with fused path. Compare: +- Final loss within 10% of Candle path (different noise patterns expected) +- Learning curve shape similar + +- [ ] **Step 5: Commit** + +```bash +git add crates/ml/src/cuda_pipeline/dqn_training_kernel.cu \ + crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs \ + crates/ml/src/cuda_pipeline/mod.rs \ + crates/ml-dqn/src/dqn.rs +git commit -m "feat(cuda): fused forward+loss+backward+adam training kernel with CUDA Graph" +``` diff --git a/docs/superpowers/specs/2026-03-15-fused-cuda-training-design.md b/docs/superpowers/specs/2026-03-15-fused-cuda-training-design.md new file mode 100644 index 000000000..084e011d1 --- /dev/null +++ b/docs/superpowers/specs/2026-03-15-fused-cuda-training-design.md @@ -0,0 +1,416 @@ +# Fused CUDA Training Kernel — Design Spec + +## Problem + +The DQN training loop dispatches **2,100+ CUDA kernels per batch** through Candle: +- Forward: ~50 dispatches (3 network passes × ~17 ops) +- Loss: ~60 dispatches (branching C51 cross-entropy) +- Backward: ~1,000 dispatches (autograd traversal) +- Gradient clipping: ~500 dispatches (per-parameter norm + scale) +- Adam optimizer: ~500 dispatches (per-parameter update) + +At 1,531 batches/epoch = **3.2M kernel launches/epoch**. Each launch carries ~5–10μs CPU overhead. On H100, the network is tiny (210K params) — pure compute per batch is <1μs, but launch overhead adds **~20ms per batch**. Total: ~30s/epoch of pure dispatch latency vs <1.5s of actual compute. + +**Root cause:** Candle's `cuMemAlloc_v2` allocations during backward invalidate CUDA Graph capture. We cannot wrap the existing Candle chain in a graph. + +## Solution + +Replace the entire Candle training dispatch chain with **3 fused CUDA kernels**, captured in a **single CUDA Graph**: + +``` +Per-batch (outside graph): PER sample → GPU batch ready +CUDA Graph replay (1 launch): [forward_loss → backward → adam_update] +Per-batch (outside graph): PER priority update → target network EMA (every N steps) +``` + +Dispatch reduction: **2,100 → 1** per batch. Epoch speedup: **~20x** on launch overhead. + +## Architecture + +### Network: Branching Dueling DQN with C51 + +Production config (from BranchingConfig): +- `state_dim`: 48 (no OFI) or 56 (with OFI), tensor-core aligned +- `shared_hidden_dims`: [256, 256] +- `value_hidden_dim`: 128, outputs [1, num_atoms] = [1, 51] +- `branch_hidden_dim`: 128 +- `branch_sizes`: [5, 3, 3] (exposure, order, urgency) +- `num_atoms`: 51 (C51 distributional) +- `v_min`: -25.0, `v_max`: 25.0 +- `use_distributional`: true +- `use_noisy`: true (NoisyNet in heads) +- Activation: LeakyReLU(0.01) after shared/FC layers + +### Weight Layout (20 tensors, ~210K params) + +From DuelingWeightSet + BranchingWeightSet in `gpu_weights.rs`: + +| Layer | Weight | Shape | Count | +|-------|--------|-------|-------| +| shared_0 | w_s1, b_s1 | [256, STATE_DIM], [256] | ~14.6K | +| shared_1 | w_s2, b_s2 | [256, 256], [256] | 65.8K | +| value_fc | w_v1, b_v1 | [128, 256], [128] | 32.9K | +| value_out | w_v2, b_v2 | [51, 128], [51] | 6.6K | +| branch_0_fc | w_a1, b_a1 | [128, 256], [128] | 32.9K | +| branch_0_out | w_a2, b_a2 | [255, 128], [255] | 32.9K | +| branch_1_fc | w_bo1, b_bo1 | [128, 256], [128] | 32.9K | +| branch_1_out | w_bo2, b_bo2 | [153, 128], [153] | 19.7K | +| branch_2_fc | w_bu1, b_bu1 | [128, 256], [128] | 32.9K | +| branch_2_out | w_bu2, b_bu2 | [153, 128], [153] | 19.7K | +| **Total** | | | **~291K** | + +Note: Branch output dims = n_d × num_atoms (distributional). Non-distributional: n_d only. + +### Three Forward Passes Per Batch + +1. **Online on states** (training=true, noise+dropout active) → current Q-distribution +2. **Target on next_states** (training=false) → target Q-distribution for Bellman +3. **Online on next_states** (training=false) → action selection for Double DQN + +### C51 Loss (Per Branch) + +For each branch d ∈ {0,1,2}: +1. Gather current log-probs for taken action aᵈ: [batch, num_atoms] +2. Get target probs for best next action (Double DQN selects): [batch, num_atoms] +3. Bellman projection: Tᵤ = r + γz (1-done), project onto support atoms via linear interpolation +4. Cross-entropy: CEᵈ = -Σⱼ projected[j] × current_log_probs[j] + +Final: loss = (1/3) Σᵈ CEᵈ × IS_weight + +## Phase 1: Forward + Loss Kernel + +**File:** `crates/ml/src/cuda_pipeline/dqn_training_kernel.cu` + +### Kernel Signature + +```c +extern "C" __global__ void dqn_forward_loss_kernel( + // Batch data (from GPU PER) + const float* states, // [B, STATE_DIM] + const float* next_states, // [B, STATE_DIM] + const int* actions, // [B] factored action indices (0-44) + const float* rewards, // [B] + const float* dones, // [B] + const float* is_weights, // [B] PER importance-sampling weights + // Online network weights (20 tensors via #define pointers) + // Target network weights (20 tensors) + // Saved activations output (for backward kernel) + float* saved_h_s1, // [B, SHARED_H1] + float* saved_h_s2, // [B, SHARED_H2] + float* saved_h_v1, // [B, VALUE_H] + float* saved_h_bd, // [B, 3, ADV_H] (3 branches) + float* saved_logits, // [B, total_branch_output_atoms] + float* saved_target_probs, // [B, 3, NUM_ATOMS] projected targets (detached) + // Outputs + float* out_loss, // [B] per-sample weighted loss + float* out_td_errors, // [B] for PER priority update + float* out_total_loss, // [1] mean loss (for monitoring) + // Config + int batch_size, + float gamma, + float v_min, float v_max, int num_atoms +); +``` + +### Thread Layout + +- grid = (batch_size, 1, 1) — one block per sample +- block = (32, 1, 1) — one warp per block (matches existing pattern) +- Shared memory: SHMEM_TILE_ROWS × SHMEM_MAX_IN_DIM for weight tiling + +### Algorithm (per block/sample) + +1. Load state[i] into distributed registers (stride-32) +2. **Online forward (training):** + - shared_0 → LeakyReLU → shared_1 → LeakyReLU → save h_s2 + - value_fc → LeakyReLU → save h_v1 → value_out → value_logits [51] + - For each branch d: branch_fc → LeakyReLU → save h_bd → branch_out → branch_logits [n_d×51] + - Compute dueling: Q_d = V + A_d - mean(A_d) per atom + - log_softmax per action within branch → save logits for backward +3. **Target forward (inference):** Same architecture with target weights, no saves +4. **Online forward on next_states (inference):** For Double DQN action selection +5. **Action decomposition:** factored_action → (exposure_idx, order_idx, urgency_idx) +6. **C51 loss per branch:** + - Gather current log-probs for taken action + - Select best next action (from online forward on next_states) + - Gather target probs for best next action + - Bellman projection: T_z = r + γ×z×(1-done), clip to [v_min, v_max] + - Linear interpolation onto support atoms (scatter-add) + - Cross-entropy: -Σ projected × current_log_probs +7. **Average over branches, multiply by IS weight** +8. **Write:** per-sample loss, td_error, atomicAdd to total_loss + +### NoisyNet Handling + +NoisyNet layers use: y = (μ_w + σ_w ⊙ ε_w) x + (μ_b + σ_b ⊙ ε_b) + +For the fused kernel, noise tensors ε are pre-generated on GPU before the graph replay. The kernel reads them as additional input buffers. This allows different noise per batch while keeping the graph structure fixed. + +## Phase 2: Backward Kernel + +```c +extern "C" __global__ void dqn_backward_kernel( + // Saved activations (from forward kernel) + const float* saved_h_s1, saved_h_s2, saved_h_v1, saved_h_bd, saved_logits, + const float* saved_target_probs, + const float* states, // [B, STATE_DIM] for input gradients + const float* is_weights, // [B] + // Online network weights (for W^T in backward) + // ... (same weight pointers as forward) + // Gradient output buffers (zero-initialized before launch) + float* grad_w_s1, float* grad_b_s1, // [256, STATE_DIM], [256] + float* grad_w_s2, float* grad_b_s2, // ... + // ... all 20 gradient tensors + int batch_size +); +``` + +### Backward Math + +For each sample i (one block per sample): + +1. **∂L/∂logits** from C51 cross-entropy: + - For each branch d, for the taken action aᵈ: + - ∂CE/∂log_p = -projected_target (shape: [num_atoms]) + - Through log_softmax: ∂L/∂z = ∂L/∂log_p - softmax(z) × Σ(∂L/∂log_p) + +2. **Backward through branch output layers** (linear, no activation): + - ∂L/∂W_out = ∂L/∂z × h_bd^T → atomicAdd to grad buffer + - ∂L/∂b_out = ∂L/∂z → atomicAdd + - ∂L/∂h_bd = W_out^T × ∂L/∂z + +3. **Backward through branch FC** (linear + LeakyReLU): + - ∂L/∂pre_bd = ∂L/∂h_bd ⊙ leaky_relu'(pre_bd) + - ∂L/∂W_fc = ∂L/∂pre_bd × h_s2^T → atomicAdd + - ∂L/∂b_fc = ∂L/∂pre_bd → atomicAdd + - Accumulate ∂L/∂h_s2 from all 3 branches + value head + +4. **Backward through value head** (similar to branches) + +5. **Backward through shared layers** (accumulated gradients from value + all branches): + - shared_1 backward + - shared_0 backward + +6. **Gradient accumulation:** Each block atomicAdds its per-sample gradient contribution to global gradient buffers. + +### atomicAdd Strategy + +- 291K gradient elements +- batch_size = 256 concurrent blocks +- Average ~1 atomicAdd per element (256 writes to 291K elements) +- H100 L2 cache (50MB) holds all gradient buffers — atomics are fast + +## Phase 3: Adam Optimizer Kernel + +```c +extern "C" __global__ void dqn_adam_kernel( + // Parameters (read + write, flattened) + float* params, // [TOTAL_PARAMS] all weights flattened + // Gradient buffer (read only) + const float* grads, // [TOTAL_PARAMS] + // Adam state (read + write) + float* m, // [TOTAL_PARAMS] first moment + float* v, // [TOTAL_PARAMS] second moment + // Config + float lr, float beta1, float beta2, float eps, + float weight_decay, float max_grad_norm, + int step, // for bias correction + int total_params, + // Output + float* out_grad_norm // [1] pre-clip gradient L2 norm +); +``` + +### Algorithm + +1. **Gradient norm** (parallel reduction over all params): + - Each thread accumulates partial norm_sq for its elements + - Block-level reduction → global atomicAdd → shared grad_norm + +2. **Gradient clipping:** + - scale = min(max_grad_norm / (grad_norm + eps), 1.0) + - Each thread: grad[i] *= scale + +3. **Adam update** (per element, trivially parallel): + - m[i] = β1 × m[i] + (1-β1) × grad[i] + - v[i] = β2 × v[i] + (1-β2) × grad[i]² + - m_hat = m[i] / (1 - β1^step) + - v_hat = v[i] / (1 - β2^step) + - params[i] -= lr × m_hat / (√v_hat + eps) + weight_decay × params[i] + +Grid: (ceil(TOTAL_PARAMS / 256), 1, 1), block: (256, 1, 1) + +## Phase 4: CUDA Graph Integration + +### Graph Structure + +``` +PER sample (outside graph) + ↓ +begin_capture(THREAD_LOCAL) + ↓ +[noise_generate_kernel] ← NoisyNet noise (if enabled) + ↓ +[dqn_forward_loss_kernel] + ↓ +[zero_grad_kernel] ← memset gradient buffers to 0 + ↓ +[dqn_backward_kernel] + ↓ +[dqn_adam_kernel] + ↓ +end_capture(AUTO_FREE_ON_LAUNCH) + ↓ +graph.launch() per batch (replay) + ↓ +PER priority update (outside graph) +Target EMA update (outside graph, every N steps) +``` + +### Input Tensor Updates Between Replays + +CUDA Graph replays with fixed buffer pointers. Between replays: +- Write new batch data (states, next_states, actions, rewards, dones, is_weights) into **pre-allocated fixed buffers** +- The PER sampling writes directly into these buffers +- No new allocations between replays + +### Graph Invalidation + +Rebuild graph when: +- batch_size changes (should not happen — fixed in config) +- Network architecture changes (never during training) +- First batch of training (initial capture) + +### Existing Pattern + +Follow `gpu_backtest_evaluator.rs`: +- `SendSyncGraph` wrapper for `CudaGraph` +- `CU_STREAM_CAPTURE_MODE_THREAD_LOCAL` +- `CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH` + +## Phase 5: Rust Host Code + +**File:** `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` + +### GpuDqnTrainer Struct + +```rust +pub struct GpuDqnTrainer { + stream: Arc, + // Compiled kernels + forward_loss_fn: CudaFunction, + backward_fn: CudaFunction, + adam_fn: CudaFunction, + // Pre-allocated buffers (fixed size for graph capture) + batch_states: CudaSlice, // [B, STATE_DIM] + batch_next_states: CudaSlice, // [B, STATE_DIM] + batch_actions: CudaSlice, // [B] + batch_rewards: CudaSlice, // [B] + batch_dones: CudaSlice, // [B] + batch_is_weights: CudaSlice, // [B] + // Activation save buffers + saved_activations: SavedActivations, + // Gradient buffers (zero before each backward) + gradient_buf: CudaSlice, // [TOTAL_PARAMS] + // Adam state + adam_m: CudaSlice, // [TOTAL_PARAMS] + adam_v: CudaSlice, // [TOTAL_PARAMS] + // Outputs + per_sample_loss: CudaSlice, // [B] + td_errors: CudaSlice, // [B] + total_loss: CudaSlice, // [1] + grad_norm: CudaSlice, // [1] + // CUDA Graph + training_graph: Option, + // NoisyNet noise buffers + noise_buffers: Option, +} +``` + +### Integration with DQN train_step() + +```rust +// In dqn.rs train_step(): +#[cfg(feature = "cuda")] +if let Some(ref mut gpu_trainer) = self.gpu_trainer { + // 1. PER samples into gpu_trainer's fixed input buffers + self.memory.sample_into(&mut gpu_trainer.batch_buffers)?; + + // 2. Generate NoisyNet noise (if enabled) + gpu_trainer.generate_noise()?; + + // 3. Replay CUDA Graph (or capture on first batch) + gpu_trainer.train_step()?; + + // 4. Update PER priorities from td_errors output + self.memory.update_priorities_gpu( + &gpu_trainer.batch_indices, &gpu_trainer.td_errors + )?; + + // 5. Target network EMA (every N steps, outside graph) + self.update_target_networks()?; + + return Ok(GpuTrainResult { + loss_gpu: gpu_trainer.total_loss_tensor()?, + grad_norm_gpu: gpu_trainer.grad_norm_tensor()?, + }); +} +``` + +## Dimension Injection + +Follow existing pattern from `gpu_experience_collector.rs` — all dimensions injected via `#define` before NVRTC compilation: + +```rust +let dim_overrides = format!( + "#define STATE_DIM {state_dim}\n\ + #define MARKET_DIM {market_dim}\n\ + #define PORTFOLIO_DIM 3\n\ + #define SHARED_H1 {shared_h1}\n\ + #define SHARED_H2 {shared_h2}\n\ + #define VALUE_H {value_h}\n\ + #define ADV_H {adv_h}\n\ + #define NUM_ATOMS {num_atoms}\n\ + #define V_MIN {v_min:.6}\n\ + #define V_MAX {v_max:.6}\n\ + #define BRANCH_0_SIZE {branch_0}\n\ + #define BRANCH_1_SIZE {branch_1}\n\ + #define BRANCH_2_SIZE {branch_2}\n\ + #define BATCH_SIZE {batch_size}\n\ + #define TOTAL_PARAMS {total_params}\n\ + #define USE_NOISY_NETS {use_noisy}\n\ + #define USE_DOUBLE_DQN 1\n" +); +``` + +## Test Strategy + +1. **Numerical correctness test:** Run fused kernel and Candle path on same input, compare: + - Forward outputs (Q-values) within 1e-5 + - Loss values within 1e-4 + - Gradients within 1e-3 (accumulated atomicAdd tolerance) + - Adam-updated params within 1e-4 + +2. **CUDA Graph smoke test:** Capture + 10 replays, verify loss decreases monotonically + +3. **Integration test:** Full training epoch with fused path, compare learning curve against Candle baseline + +## Risk Assessment + +| Risk | Mitigation | +|------|-----------| +| Gradient correctness | Numerical diff test against Candle autograd | +| atomicAdd precision | Use float atomicAdd (sufficient for 256-sample batch) | +| NoisyNet noise patterns | Pre-generated noise buffers, updated between graph replays | +| Graph invalidation | Only rebuild on first batch (all dims fixed) | +| Occupancy concerns | Network is tiny — compute <1μs, bottleneck is purely launch overhead | +| Dropout in training | Compile-time disabled for fused path (NoisyNet replaces dropout for exploration) | + +## Performance Targets + +| Metric | Current | Target | +|--------|---------|--------| +| Kernel launches/batch | 2,100 | 1 (graph replay) | +| Launch overhead/batch | ~15ms | ~1μs | +| Epoch time (launch only) | ~23s | ~1.5ms | +| Expected epoch speedup | - | ~15-20x | diff --git a/scripts/gpu-hotpath-guard.sh b/scripts/gpu-hotpath-guard.sh index d2ae99582..4573281f3 100755 --- a/scripts/gpu-hotpath-guard.sh +++ b/scripts/gpu-hotpath-guard.sh @@ -33,7 +33,6 @@ HOT_PATHS=( "gpu_backtest_evaluator" "gpu_training_guard" "gpu_weights" - "mixed_precision.rs" # DQN crate — all modules (inference, networks, replay, regularization) "ml-dqn/src/" # PPO training + inference diff --git a/services/ml_training_service/src/ensemble_training_coordinator.rs b/services/ml_training_service/src/ensemble_training_coordinator.rs index f3c5ab04e..f5137c8af 100644 --- a/services/ml_training_service/src/ensemble_training_coordinator.rs +++ b/services/ml_training_service/src/ensemble_training_coordinator.rs @@ -603,7 +603,6 @@ mod tests { performance_config: PerformanceConfig { device_preference: "cpu".to_string(), max_memory_bytes: 4_000_000_000, - mixed_precision: false, num_workers: 2, gradient_accumulation_steps: 1, }, diff --git a/services/ml_training_service/src/gpu_config.rs b/services/ml_training_service/src/gpu_config.rs index cf1253f38..6d52218c2 100644 --- a/services/ml_training_service/src/gpu_config.rs +++ b/services/ml_training_service/src/gpu_config.rs @@ -16,8 +16,6 @@ pub struct GpuConfig { pub device_id: u32, /// Maximum GPU memory to use in GB pub max_memory_gb: f32, - /// Enable mixed precision training - pub enable_mixed_precision: bool, /// Enable GPU memory optimization pub enable_memory_optimization: bool, /// Batch size optimization factor @@ -33,7 +31,6 @@ impl Default for GpuConfig { Self { device_id: 0, max_memory_gb: 8.0, - enable_mixed_precision: true, enable_memory_optimization: true, batch_size_factor: 1.0, enable_cuda_graphs: false, @@ -114,10 +111,6 @@ impl GpuConfigManager { .and_then(|v| v.as_f64().map(|n| n as f32)) .unwrap_or(8.0); - let enable_mixed_precision = get_value("gpu_enable_mixed_precision") - .and_then(|v| v.as_bool()) - .unwrap_or(true); - let enable_memory_optimization = get_value("gpu_enable_memory_optimization") .and_then(|v| v.as_bool()) .unwrap_or(true); @@ -137,7 +130,6 @@ impl GpuConfigManager { Ok(GpuConfig { device_id, max_memory_gb, - enable_mixed_precision, enable_memory_optimization, batch_size_factor, enable_cuda_graphs, @@ -236,14 +228,6 @@ impl GpuConfigManager { } } - /// Check if mixed precision is enabled - pub fn is_mixed_precision_enabled(&self) -> bool { - self.gpu_config - .as_ref() - .map(|c| c.enable_mixed_precision) - .unwrap_or(false) - } - /// Check if memory optimization is enabled pub fn is_memory_optimization_enabled(&self) -> bool { self.gpu_config @@ -279,7 +263,6 @@ mod tests { let config = GpuConfig::default(); assert_eq!(config.device_id, 0); assert_eq!(config.max_memory_gb, 8.0); - assert!(config.enable_mixed_precision); assert!(config.enable_memory_optimization); assert_eq!(config.batch_size_factor, 1.0); assert!(!config.enable_cuda_graphs); diff --git a/services/ml_training_service/tests/ensemble_training_tests.rs b/services/ml_training_service/tests/ensemble_training_tests.rs index fca3622b6..1a42c6c35 100644 --- a/services/ml_training_service/tests/ensemble_training_tests.rs +++ b/services/ml_training_service/tests/ensemble_training_tests.rs @@ -513,7 +513,6 @@ fn create_model_config( performance_config: PerformanceConfig { device_preference: "cpu".to_string(), max_memory_bytes: 4_000_000_000, - mixed_precision: false, num_workers: 2, gradient_accumulation_steps: 1, }, diff --git a/services/ml_training_service/tests/fixtures/job_configs/invalid_zero_epochs.json b/services/ml_training_service/tests/fixtures/job_configs/invalid_zero_epochs.json index 009d27d29..17d9a466d 100644 --- a/services/ml_training_service/tests/fixtures/job_configs/invalid_zero_epochs.json +++ b/services/ml_training_service/tests/fixtures/job_configs/invalid_zero_epochs.json @@ -22,7 +22,6 @@ "device_preference": "cpu", "num_workers": 2, "max_memory_bytes": 2147483648, - "mixed_precision": false, "gradient_accumulation_steps": 1 } } diff --git a/services/ml_training_service/tests/fixtures/job_configs/valid_dqn.json b/services/ml_training_service/tests/fixtures/job_configs/valid_dqn.json index 609cee872..7b51ebb30 100644 --- a/services/ml_training_service/tests/fixtures/job_configs/valid_dqn.json +++ b/services/ml_training_service/tests/fixtures/job_configs/valid_dqn.json @@ -22,7 +22,6 @@ "device_preference": "cpu", "num_workers": 4, "max_memory_bytes": 8589934592, - "mixed_precision": false, "gradient_accumulation_steps": 1 } } diff --git a/services/ml_training_service/tests/integration/end_to_end_batch_workflow_test.rs b/services/ml_training_service/tests/integration/end_to_end_batch_workflow_test.rs index 9cb80760b..b08af8326 100644 --- a/services/ml_training_service/tests/integration/end_to_end_batch_workflow_test.rs +++ b/services/ml_training_service/tests/integration/end_to_end_batch_workflow_test.rs @@ -62,7 +62,6 @@ fn create_test_training_config() -> ProductionTrainingConfig { device_preference: "cpu".to_string(), num_workers: 2, max_memory_bytes: 2 * 1024 * 1024 * 1024, // 2GB - mixed_precision: false, gradient_accumulation_steps: 1, }, } diff --git a/services/ml_training_service/tests/integration/failure_recovery_test.rs b/services/ml_training_service/tests/integration/failure_recovery_test.rs index 201033a59..f4b3ae236 100644 --- a/services/ml_training_service/tests/integration/failure_recovery_test.rs +++ b/services/ml_training_service/tests/integration/failure_recovery_test.rs @@ -117,7 +117,6 @@ fn create_test_config() -> ml::training_pipeline::ProductionTrainingConfig { device_preference: "cpu".to_string(), num_workers: 2, max_memory_bytes: 2 * 1024 * 1024 * 1024, - mixed_precision: false, gradient_accumulation_steps: 1, }, } diff --git a/services/ml_training_service/tests/integration/grpc_api_integration_test.rs b/services/ml_training_service/tests/integration/grpc_api_integration_test.rs index 75dde54c8..918bf36eb 100644 --- a/services/ml_training_service/tests/integration/grpc_api_integration_test.rs +++ b/services/ml_training_service/tests/integration/grpc_api_integration_test.rs @@ -117,7 +117,6 @@ fn create_test_config() -> ml::training_pipeline::ProductionTrainingConfig { device_preference: "cpu".to_string(), num_workers: 2, max_memory_bytes: 2 * 1024 * 1024 * 1024, - mixed_precision: false, gradient_accumulation_steps: 1, }, } diff --git a/services/ml_training_service/tests/integration/real_data_integration_test.rs b/services/ml_training_service/tests/integration/real_data_integration_test.rs index a6de33519..fe1d251a6 100644 --- a/services/ml_training_service/tests/integration/real_data_integration_test.rs +++ b/services/ml_training_service/tests/integration/real_data_integration_test.rs @@ -191,7 +191,6 @@ async fn test_train_single_model_real_data() { device_preference: "cpu".to_string(), num_workers: 1, max_memory_bytes: 1024 * 1024 * 1024, // 1GB - mixed_precision: false, gradient_accumulation_steps: 1, }, }; @@ -276,7 +275,6 @@ async fn test_train_all_models_real_data() { device_preference: "cpu".to_string(), num_workers: 1, max_memory_bytes: 1024 * 1024 * 1024, - mixed_precision: false, gradient_accumulation_steps: 1, }, }; diff --git a/services/ml_training_service/tests/integration_tests.rs b/services/ml_training_service/tests/integration_tests.rs index 183e7149b..eb448584e 100644 --- a/services/ml_training_service/tests/integration_tests.rs +++ b/services/ml_training_service/tests/integration_tests.rs @@ -71,7 +71,6 @@ fn create_test_training_config() -> ProductionTrainingConfig { device_preference: "cpu".to_string(), num_workers: 4, max_memory_bytes: 8 * 1024 * 1024 * 1024, // 8GB - mixed_precision: false, gradient_accumulation_steps: 1, }, } diff --git a/services/ml_training_service/tests/orchestrator_comprehensive_tests.rs b/services/ml_training_service/tests/orchestrator_comprehensive_tests.rs index 3ece5f87b..18f0b4e6b 100644 --- a/services/ml_training_service/tests/orchestrator_comprehensive_tests.rs +++ b/services/ml_training_service/tests/orchestrator_comprehensive_tests.rs @@ -57,7 +57,6 @@ fn create_test_training_config() -> ProductionTrainingConfig { performance_config: PerformanceConfig { device_preference: "cpu".to_string(), max_memory_bytes: 8 * 1024 * 1024 * 1024, // 8GB - mixed_precision: false, num_workers: 4, gradient_accumulation_steps: 1, }, diff --git a/services/ml_training_service/tests/training_error_recovery_tests.rs b/services/ml_training_service/tests/training_error_recovery_tests.rs index c5c4b7b00..3750f58a6 100644 --- a/services/ml_training_service/tests/training_error_recovery_tests.rs +++ b/services/ml_training_service/tests/training_error_recovery_tests.rs @@ -59,7 +59,6 @@ fn create_minimal_config() -> ProductionTrainingConfig { performance_config: PerformanceConfig { device_preference: "cpu".to_string(), max_memory_bytes: 1024 * 1024 * 1024, // 1GB - mixed_precision: false, num_workers: 1, gradient_accumulation_steps: 1, }, diff --git a/services/ml_training_service/tests/validation_pipeline_tests.rs b/services/ml_training_service/tests/validation_pipeline_tests.rs index b00edcf97..387bdee05 100644 --- a/services/ml_training_service/tests/validation_pipeline_tests.rs +++ b/services/ml_training_service/tests/validation_pipeline_tests.rs @@ -435,7 +435,6 @@ fn create_completed_training_job(job_id: Uuid) -> TrainingJob { performance_config: PerformanceConfig { device_preference: "cpu".to_string(), max_memory_bytes: 8 * 1024 * 1024 * 1024, - mixed_precision: false, num_workers: 4, gradient_accumulation_steps: 1, }, diff --git a/services/trading_service/src/services/ppo_model.rs b/services/trading_service/src/services/ppo_model.rs index 11365af79..dfc54908e 100644 --- a/services/trading_service/src/services/ppo_model.rs +++ b/services/trading_service/src/services/ppo_model.rs @@ -170,7 +170,6 @@ impl PPOModel { lstm_sequence_length: 32, accumulation_steps: 1, clip_epsilon_high: None, - mixed_precision: None, use_symlog: true, use_adaptive_entropy: true, use_percentile_scaling: true, diff --git a/services/trading_service/src/services/tft_model.rs b/services/trading_service/src/services/tft_model.rs index 94e36242e..701cfef10 100644 --- a/services/trading_service/src/services/tft_model.rs +++ b/services/trading_service/src/services/tft_model.rs @@ -50,7 +50,6 @@ impl TFTModel { dropout_rate: 0.1, l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: false, memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000, @@ -173,7 +172,6 @@ mod tests { dropout_rate: 0.1, l2_regularization: 1e-4, use_flash_attention: true, - mixed_precision: false, memory_efficient: true, max_inference_latency_us: 50, target_throughput_pps: 100_000,