refactor(ml): delete mixed_precision module — BF16 unconditional on CUDA

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 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-16 16:11:48 +01:00
parent cd54a6f27d
commit d95e205d4b
132 changed files with 1139 additions and 1861 deletions

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -127,7 +127,6 @@ pub struct HardwareConfig {
pub use_gpu: bool,
pub gpu_memory_limit_mb: Option<usize>,
pub cpu_threads: Option<usize>,
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
}
}
}

View File

@@ -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;

View File

@@ -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<MixedPrecisionConfig> {
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<MixedPrecisionConfig> {
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<candle_core::Tensor, candle_core::Error> {
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, MLError> {
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, MLError> {
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<F>(
input: &Tensor,
config: &MixedPrecisionConfig,
forward_fn: F,
) -> Result<Tensor, MLError>
where
F: Fn(&Tensor) -> Result<Tensor, MLError>,
{
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<Tensor, MLError> {
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<Tensor, MLError> {
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::<f32>()
.map_err(|e| MLError::TensorOperationError(e.to_string()))?;
let restored_data = restored
.to_vec1::<f32>()
.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<Tensor, MLError> {
// 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::<f32>()
.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<Tensor, MLError> {
// 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::<f32>()
.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<Tensor, MLError> {
// 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::<f32>()
.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::<f32>()
.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::<f32>()
.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::<f32>()
.map_err(|e| MLError::TensorOperationError(e.to_string()))?;
let unscaled_data = unscaled
.to_vec1::<f32>()
.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<Tensor, MLError> {
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<Tensor, MLError> {
// 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::<f32>()
.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);
}
}
}

View File

@@ -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<Tensor, MLError> {
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<Tensor, MLError> {
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> {

View File

@@ -255,7 +255,7 @@ impl MultiHeadAttention {
/// 5. Optional: Add residual connection and layer normalization
pub fn forward(&self, x: &Tensor, mask: Option<&Tensor>) -> Result<Tensor, MLError> {
// 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)?;

View File

@@ -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<usize>,
) -> 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<BranchOutput, MLError> {
// 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)
}

View File

@@ -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<Self, MLError> {
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)

View File

@@ -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<Self, MLError> {
// 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<Tensor, MLError> {
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)

View File

@@ -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<super::mixed_precision::MixedPrecisionConfig>,
}
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<Tensor, MLError> {
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<Tensor, MLError> {
// 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<Tensor, MLError> {
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::<Tensor>, None::<Tensor>)
}
} 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::<Tensor>, None::<Tensor>)
}
} else {

View File

@@ -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<Self, MLError> {
// 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<Tensor, MLError> {
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| {

View File

@@ -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)?;

View File

@@ -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;

View File

@@ -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<crate::mixed_precision::MixedPrecisionConfig>,
}
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<crate::mixed_precision::MixedPrecisionConfig>,
) -> CandleResult<Tensor> {
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<Tensor> {
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::<f32>().map_err(|e| {

View File

@@ -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)?;

View File

@@ -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<Self, MLError> {
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<Tensor, MLError> {
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)?;

View File

@@ -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

View File

@@ -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;

View File

@@ -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<Self, MLError> {
// 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<Self, MLError> {
// SegmentTree: 2 * next_power_of_two(capacity) f32 entries
// PrioritizedReplayBuffer: capacity Option<Experience> 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::<f32>());
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<BatchSample, MLError> {
match self {

View File

@@ -108,7 +108,7 @@ impl ResidualBlock {
/// Output tensor with same shape as input
pub fn forward(&self, x: &Tensor, train: bool) -> Result<Tensor, MLError> {
// 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

View File

@@ -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)?;

View File

@@ -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,
}
}

View File

@@ -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();

View File

@@ -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::<f32>()?;
@@ -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)?;

View File

@@ -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<Self, MLError> {
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)))?,
);

View File

@@ -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<ContinuousTrajectoryTensors, MLError> {
let batch_size = self.states.len();
let dtype = training_dtype(device);
let dtype = candle_core::DType::BF16;
let state_flat: Vec<f32> = 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<ContinuousTrajectoryTensors, MLError> {
let batch_size = self.states.len();
let dtype = training_dtype(device);
let dtype = candle_core::DType::BF16;
let state_flat: Vec<f32> = 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(),

View File

@@ -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<Self, 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);
// 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<Tensor, MLError> {
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, MLError> {
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)))
}

View File

@@ -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<Self, MLError> {
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)?;

View File

@@ -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<Self, MLError> {
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<Self, MLError> {
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

View File

@@ -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<f32>,
/// 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<ml_core::mixed_precision::MixedPrecisionConfig>,
/// 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<Linear>,
device: Device,
vars: VarMap,
mixed_precision: Option<ml_core::mixed_precision::MixedPrecisionConfig>,
}
impl PolicyNetwork {
@@ -311,7 +305,7 @@ impl PolicyNetwork {
device: Device,
) -> Result<Self, MLError> {
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<Tensor, MLError> {
self.forward_mixed(input, &self.mixed_precision)
}
/// Set mixed precision config for this network
pub const fn set_mixed_precision(&mut self, mp: Option<ml_core::mixed_precision::MixedPrecisionConfig>) {
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<ml_core::mixed_precision::MixedPrecisionConfig>,
) -> Result<Tensor, MLError> {
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<Linear>,
device: Device,
vars: VarMap,
mixed_precision: Option<ml_core::mixed_precision::MixedPrecisionConfig>,
}
impl ValueNetwork {
/// Create new value network
pub fn new(input_dim: usize, hidden_dims: &[usize], device: Device) -> Result<Self, MLError> {
let vars = VarMap::new();
let var_builder = VarBuilder::from_varmap(&vars, 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<Tensor, MLError> {
self.forward_mixed(input, &self.mixed_precision)
}
/// Set mixed precision config for this network
pub const fn set_mixed_precision(&mut self, mp: Option<ml_core::mixed_precision::MixedPrecisionConfig>) {
self.mixed_precision = mp;
}
/// Forward pass with optional mixed precision (BF16/FP16).
pub fn forward_mixed(
&self,
input: &Tensor,
mixed_precision: &Option<ml_core::mixed_precision::MixedPrecisionConfig>,
) -> Result<Tensor, MLError> {
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 {}: {}",

View File

@@ -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<TrajectoryTensors, MLError> {
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<TrajectoryTensors, MLError> {
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);

View File

@@ -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<Tensor, MLError> {
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();

View File

@@ -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);

View File

@@ -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();

View File

@@ -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<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::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();

View File

@@ -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;

View File

@@ -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<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 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,

View File

@@ -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<Self, MLError> {
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)?;

View File

@@ -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<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::ModelError(e.to_string()))?;
let start = Instant::now();

View File

@@ -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)]

View File

@@ -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<Tensor, MLError> {
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),

View File

@@ -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

View File

@@ -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"))?;

View File

@@ -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();

View File

@@ -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"))?;

View File

@@ -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();
}

View File

@@ -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);
}

View File

@@ -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();

View File

@@ -79,7 +79,7 @@ impl XLSTMNetwork {
///
/// Returns `(batch, output_dim)`.
pub fn forward(&self, input: &Tensor) -> Result<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::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());
}

View File

@@ -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();

View File

@@ -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()
};

View File

@@ -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]

View File

@@ -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<()> {

View File

@@ -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")

View File

@@ -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 {

View File

@@ -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<Self, MLError> {
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<Tensor, MLError> {
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());

View File

@@ -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

View File

@@ -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]

View File

@@ -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

View File

@@ -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

View File

@@ -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)

View File

@@ -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]

View File

@@ -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) ---

View File

@@ -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<Tensor> {
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<Self> {
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<Self> {
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

View File

@@ -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<Tensor> {
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

View File

@@ -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

View File

@@ -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

View File

@@ -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)

View File

@@ -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();

View File

@@ -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,

View File

@@ -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,

View File

@@ -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,

View File

@@ -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<Self, MLError> {
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<Tensor, MLError> {
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)
}

View File

@@ -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;

View File

@@ -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)?;

View File

@@ -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<Self> {
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(

View File

@@ -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

View File

@@ -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]

View File

@@ -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<Tensor, MLError> {
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

View File

@@ -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<Tensor, MLError> {
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

View File

@@ -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<Vec<Experience>, 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<f32> = states.to_dtype(DType::F32)
.and_then(|t| t.to_vec2::<f32>())
.map_err(|e| MLError::TrainingError(format!("states download: {e}")))?
.into_iter().flatten().collect();
let next_cpu: Vec<f32> = next_states.to_dtype(DType::F32)
.and_then(|t| t.to_vec2::<f32>())
.map_err(|e| MLError::TrainingError(format!("next_states download: {e}")))?
.into_iter().flatten().collect();
let actions_cpu: Vec<u32> = actions.to_dtype(DType::U32)
.and_then(|t| t.to_vec1::<u32>())
.map_err(|e| MLError::TrainingError(format!("actions download: {e}")))?;
let rewards_cpu: Vec<f32> = rewards.to_dtype(DType::F32)
.and_then(|t| t.to_vec1::<f32>())
.map_err(|e| MLError::TrainingError(format!("rewards download: {e}")))?;
let dones_cpu: Vec<f32> = dones.to_dtype(DType::F32)
.and_then(|t| t.to_vec1::<f32>())
.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<usize>,
/// Mixed precision for GPU training (auto-detected from HardwareBudget)
pub mixed_precision: Option<crate::dqn::mixed_precision::MixedPrecisionConfig>,
/// 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()
}
}

View File

@@ -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<f64> {
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<Tensor> = 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)
}

View File

@@ -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.

View File

@@ -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();

View File

@@ -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;

View File

@@ -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<f32> = 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.

View File

@@ -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

View File

@@ -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<f32>.
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<f32> = 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<f32> = 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<f32>.
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<f32> = 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

View File

@@ -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<f32> = if aligned > raw_dim {
let mut v = state_vec.to_vec();
v.resize(aligned, 0.0);

View File

@@ -55,10 +55,7 @@ fn create_test_trainer_with(params: DQNHyperparameters) -> Result<DQNTrainer> {
/// 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<f64>)> = vec![];

View File

@@ -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;

View File

@@ -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 {

View File

@@ -600,7 +600,7 @@ mod tests {
fn tiny_var_map() -> Result<VarMap, MLError> {
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)

View File

@@ -191,7 +191,6 @@ impl From<PpoHyperparameters> 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())?;

View File

@@ -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,

View File

@@ -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)?;

View File

@@ -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<f64>,
}
@@ -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),
}
}

View File

@@ -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,
},

Some files were not shown because too many files have changed in this diff Show More