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