feat(liquid): add LiquidTrainableAdapter implementing UnifiedTrainable
Full adapter bridging CandleCfCNetwork to the unified training pipeline with VarMap-based checkpointing, AdamW optimizer, and gradient norm tracking. Includes 10 unit tests covering creation, training steps, validation, metrics, learning rate, checkpoint roundtrip, and error cases. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
641
ml/src/liquid/adapter.rs
Normal file
641
ml/src/liquid/adapter.rs
Normal file
@@ -0,0 +1,641 @@
|
||||
//! UnifiedTrainable Adapter for Liquid CfC v2
|
||||
//!
|
||||
//! Bridges the `CandleCfCNetwork` to the unified training pipeline used by DQN/PPO/TFT/Mamba2.
|
||||
//! This adapter manages the VarMap, optimizer, and gradient lifecycle so the CfC network
|
||||
//! can participate in the standardized training orchestration.
|
||||
|
||||
use candle_core::backprop::GradStore;
|
||||
use candle_core::{Device, Tensor};
|
||||
use candle_nn::{AdamW, Optimizer, ParamsAdamW, VarBuilder, VarMap};
|
||||
use std::collections::HashMap;
|
||||
|
||||
use super::candle_cfc::{CandleCfCNetwork, CfCTrainConfig};
|
||||
use crate::training::unified_trainer::{
|
||||
checkpoint, CheckpointMetadata, TrainingMetrics, UnifiedTrainable,
|
||||
};
|
||||
use crate::MLError;
|
||||
|
||||
/// Adapter wrapping `CandleCfCNetwork` to implement `UnifiedTrainable`.
|
||||
///
|
||||
/// Owns the VarMap and optimizer so that the training loop can call
|
||||
/// `forward` / `backward` / `optimizer_step` in the standard sequence.
|
||||
pub struct LiquidTrainableAdapter {
|
||||
network: CandleCfCNetwork,
|
||||
varmap: VarMap,
|
||||
optimizer: AdamW,
|
||||
device: Device,
|
||||
step: usize,
|
||||
config: CfCTrainConfig,
|
||||
latest_metrics: TrainingMetrics,
|
||||
last_grads: Option<GradStore>,
|
||||
learning_rate: f64,
|
||||
loss_history: Vec<f64>,
|
||||
last_grad_norm: f64,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for LiquidTrainableAdapter {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("LiquidTrainableAdapter")
|
||||
.field("config", &self.config)
|
||||
.field("device", &format!("{:?}", self.device))
|
||||
.field("step", &self.step)
|
||||
.field("learning_rate", &self.learning_rate)
|
||||
.field("loss_history_len", &self.loss_history.len())
|
||||
.field("last_grad_norm", &self.last_grad_norm)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl LiquidTrainableAdapter {
|
||||
/// Create a new Liquid CfC trainable adapter.
|
||||
///
|
||||
/// Initialises the VarMap, builds the network, and creates an AdamW optimizer
|
||||
/// over all trainable parameters.
|
||||
pub fn new(config: CfCTrainConfig) -> Result<Self, MLError> {
|
||||
let device = config.device.resolve()?;
|
||||
let learning_rate = config.learning_rate;
|
||||
let varmap = VarMap::new();
|
||||
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
|
||||
|
||||
let network = CandleCfCNetwork::new(&config, &vb)?;
|
||||
|
||||
let optimizer = AdamW::new(
|
||||
varmap.all_vars(),
|
||||
ParamsAdamW {
|
||||
lr: learning_rate,
|
||||
beta1: 0.9,
|
||||
beta2: 0.999,
|
||||
eps: 1e-8,
|
||||
weight_decay: 0.0,
|
||||
},
|
||||
)
|
||||
.map_err(|e| {
|
||||
MLError::ModelError(format!("Failed to initialise AdamW optimizer: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(Self {
|
||||
network,
|
||||
varmap,
|
||||
optimizer,
|
||||
device,
|
||||
step: 0,
|
||||
config,
|
||||
latest_metrics: TrainingMetrics::default(),
|
||||
last_grads: None,
|
||||
learning_rate,
|
||||
loss_history: Vec::new(),
|
||||
last_grad_norm: 0.0,
|
||||
})
|
||||
}
|
||||
|
||||
/// Access the underlying CfC network (read-only).
|
||||
pub fn network(&self) -> &CandleCfCNetwork {
|
||||
&self.network
|
||||
}
|
||||
|
||||
/// Access the VarMap (read-only).
|
||||
pub fn varmap(&self) -> &VarMap {
|
||||
&self.varmap
|
||||
}
|
||||
}
|
||||
|
||||
impl UnifiedTrainable for LiquidTrainableAdapter {
|
||||
fn model_type(&self) -> &str {
|
||||
"Liquid-CfC"
|
||||
}
|
||||
|
||||
fn device(&self) -> &Device {
|
||||
&self.device
|
||||
}
|
||||
|
||||
/// Forward pass: expects 3D input `[batch, seq_len, features]`, returns `[batch, output_size]`.
|
||||
fn forward(&mut self, input: &Tensor) -> Result<Tensor, MLError> {
|
||||
self.network.forward(input)
|
||||
}
|
||||
|
||||
/// MSE loss between predictions and targets.
|
||||
fn compute_loss(&self, predictions: &Tensor, targets: &Tensor) -> Result<Tensor, MLError> {
|
||||
let diff = predictions.sub(targets).map_err(|e| {
|
||||
MLError::TrainingError(format!("compute_loss sub: {}", e))
|
||||
})?;
|
||||
let squared = diff.powf(2.0).map_err(|e| {
|
||||
MLError::TrainingError(format!("compute_loss powf: {}", e))
|
||||
})?;
|
||||
let loss = squared.mean_all().map_err(|e| {
|
||||
MLError::TrainingError(format!("compute_loss mean: {}", e))
|
||||
})?;
|
||||
Ok(loss)
|
||||
}
|
||||
|
||||
/// Backward pass: computes gradients and returns the L2 gradient norm.
|
||||
fn backward(&mut self, loss: &Tensor) -> Result<f64, MLError> {
|
||||
let grads = loss
|
||||
.backward()
|
||||
.map_err(|e| MLError::TrainingError(format!("Backward pass failed: {}", e)))?;
|
||||
|
||||
// Compute L2 gradient norm across all parameters
|
||||
let mut total_norm_sq = 0.0_f64;
|
||||
|
||||
let varmap_data = self
|
||||
.varmap
|
||||
.data()
|
||||
.lock()
|
||||
.map_err(|e| MLError::LockError(format!("Failed to lock VarMap: {}", e)))?;
|
||||
|
||||
for (_name, var) in varmap_data.iter() {
|
||||
if let Some(grad) = grads.get(var.as_tensor()) {
|
||||
let grad_norm_sq = grad
|
||||
.sqr()
|
||||
.and_then(|t| t.sum_all())
|
||||
.and_then(|t| t.to_dtype(candle_core::DType::F64))
|
||||
.and_then(|t| t.to_scalar::<f64>())
|
||||
.map_err(|e| {
|
||||
MLError::TrainingError(format!("Failed to compute grad norm: {}", e))
|
||||
})?;
|
||||
total_norm_sq += grad_norm_sq;
|
||||
}
|
||||
}
|
||||
|
||||
let grad_norm = total_norm_sq.sqrt();
|
||||
|
||||
if grad_norm.is_nan() || grad_norm.is_infinite() {
|
||||
return Err(MLError::TrainingError(
|
||||
"Gradient norm is NaN or Inf -- gradient explosion detected".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
self.last_grad_norm = grad_norm;
|
||||
self.latest_metrics.grad_norm = Some(grad_norm);
|
||||
|
||||
// Record loss value for history
|
||||
if let Ok(loss_val) = loss
|
||||
.to_dtype(candle_core::DType::F64)
|
||||
.and_then(|t| t.to_scalar::<f64>())
|
||||
{
|
||||
self.loss_history.push(loss_val);
|
||||
self.latest_metrics.loss = loss_val;
|
||||
}
|
||||
|
||||
// Store grads for optimizer_step (must drop lock first)
|
||||
drop(varmap_data);
|
||||
self.last_grads = Some(grads);
|
||||
|
||||
Ok(grad_norm)
|
||||
}
|
||||
|
||||
/// Apply optimizer update using the gradients from the last `backward` call.
|
||||
fn optimizer_step(&mut self) -> Result<(), MLError> {
|
||||
let grads = self.last_grads.as_ref().ok_or_else(|| {
|
||||
MLError::TrainingError(
|
||||
"No gradients available. Call backward() before optimizer_step()".to_string(),
|
||||
)
|
||||
})?;
|
||||
|
||||
self.optimizer
|
||||
.step(grads)
|
||||
.map_err(|e| MLError::TrainingError(format!("Optimizer step failed: {}", e)))?;
|
||||
|
||||
self.step += 1;
|
||||
self.last_grads = None;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Clear stored gradients.
|
||||
fn zero_grad(&mut self) -> Result<(), MLError> {
|
||||
self.last_grads = None;
|
||||
self.last_grad_norm = 0.0;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn get_learning_rate(&self) -> f64 {
|
||||
self.learning_rate
|
||||
}
|
||||
|
||||
/// Set learning rate by updating the internal optimizer.
|
||||
fn set_learning_rate(&mut self, lr: f64) -> Result<(), MLError> {
|
||||
if lr <= 0.0 {
|
||||
return Err(MLError::ValidationError {
|
||||
message: format!("Learning rate must be positive, got {}", lr),
|
||||
});
|
||||
}
|
||||
self.learning_rate = lr;
|
||||
self.optimizer.set_learning_rate(lr);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn get_step(&self) -> usize {
|
||||
self.step
|
||||
}
|
||||
|
||||
fn collect_metrics(&self) -> TrainingMetrics {
|
||||
let mut custom_metrics = HashMap::new();
|
||||
custom_metrics.insert("param_count".to_string(), self.network.param_count() as f64);
|
||||
custom_metrics.insert("step_count".to_string(), self.step as f64);
|
||||
custom_metrics.insert("last_grad_norm".to_string(), self.last_grad_norm);
|
||||
|
||||
// Number of parameters from VarMap
|
||||
if let Ok(data) = self.varmap.data().lock() {
|
||||
let num_params: usize = data
|
||||
.iter()
|
||||
.map(|(_, var)| var.as_tensor().elem_count())
|
||||
.sum();
|
||||
custom_metrics.insert("num_parameters".to_string(), num_params as f64);
|
||||
}
|
||||
|
||||
TrainingMetrics {
|
||||
loss: self.loss_history.last().copied().unwrap_or(0.0),
|
||||
val_loss: self.latest_metrics.val_loss,
|
||||
accuracy: None, // CfC is a regression model
|
||||
learning_rate: self.learning_rate,
|
||||
grad_norm: Some(self.last_grad_norm),
|
||||
custom_metrics,
|
||||
}
|
||||
}
|
||||
|
||||
/// Save model weights (safetensors) and metadata (JSON).
|
||||
fn save_checkpoint(&self, checkpoint_path: &str) -> Result<String, MLError> {
|
||||
let metadata = CheckpointMetadata {
|
||||
model_type: "Liquid-CfC".to_string(),
|
||||
version: "2.0.0".to_string(),
|
||||
epoch: 0,
|
||||
step: self.step,
|
||||
timestamp: std::time::SystemTime::now(),
|
||||
config: serde_json::to_value(&self.config).map_err(|e| {
|
||||
MLError::SerializationError {
|
||||
reason: format!("Failed to serialise config: {}", e),
|
||||
}
|
||||
})?,
|
||||
metrics: self.collect_metrics(),
|
||||
};
|
||||
|
||||
checkpoint::save_metadata(&metadata, checkpoint_path)?;
|
||||
|
||||
// Extract tensors from VarMap and save as safetensors
|
||||
let safetensors_path = format!("{}.safetensors", checkpoint_path);
|
||||
|
||||
let vars_data = self
|
||||
.varmap
|
||||
.data()
|
||||
.lock()
|
||||
.map_err(|e| MLError::LockError(format!("Failed to lock VarMap for save: {}", e)))?;
|
||||
|
||||
let mut tensors: HashMap<String, Tensor> = HashMap::new();
|
||||
for (name, var) in vars_data.iter() {
|
||||
tensors.insert(name.clone(), var.as_tensor().clone());
|
||||
}
|
||||
|
||||
candle_core::safetensors::save(&tensors, &safetensors_path)
|
||||
.map_err(|e| MLError::CheckpointError(format!("Failed to save safetensors: {}", e)))?;
|
||||
|
||||
tracing::info!(
|
||||
"Saved Liquid-CfC checkpoint to {} (step {}, {} tensors)",
|
||||
checkpoint_path,
|
||||
self.step,
|
||||
tensors.len()
|
||||
);
|
||||
|
||||
Ok(checkpoint_path.to_string())
|
||||
}
|
||||
|
||||
/// Load model weights from safetensors and restore training state from metadata.
|
||||
fn load_checkpoint(&mut self, checkpoint_path: &str) -> Result<CheckpointMetadata, MLError> {
|
||||
let metadata = checkpoint::load_metadata(checkpoint_path)?;
|
||||
|
||||
if metadata.model_type != "Liquid-CfC" {
|
||||
return Err(MLError::CheckpointError(format!(
|
||||
"Invalid model type in checkpoint: expected Liquid-CfC, got {}",
|
||||
metadata.model_type
|
||||
)));
|
||||
}
|
||||
|
||||
let safetensors_path = format!("{}.safetensors", checkpoint_path);
|
||||
let tensors = candle_core::safetensors::load(&safetensors_path, &self.device)
|
||||
.map_err(|e| {
|
||||
MLError::CheckpointError(format!("Failed to load safetensors: {}", e))
|
||||
})?;
|
||||
|
||||
let vars_data = self
|
||||
.varmap
|
||||
.data()
|
||||
.lock()
|
||||
.map_err(|e| {
|
||||
MLError::LockError(format!("Failed to lock VarMap for load: {}", e))
|
||||
})?;
|
||||
|
||||
for (name, tensor) in &tensors {
|
||||
if let Some(var) = vars_data.get(name) {
|
||||
var.set(tensor).map_err(|e| {
|
||||
MLError::CheckpointError(format!("Failed to set var {}: {}", name, e))
|
||||
})?;
|
||||
} else {
|
||||
tracing::warn!("Checkpoint contains unknown variable: {}", name);
|
||||
}
|
||||
}
|
||||
|
||||
// Restore training state
|
||||
self.step = metadata.step;
|
||||
self.latest_metrics = metadata.metrics.clone();
|
||||
|
||||
tracing::info!(
|
||||
"Loaded Liquid-CfC checkpoint from {} (step {})",
|
||||
checkpoint_path,
|
||||
metadata.step
|
||||
);
|
||||
|
||||
Ok(metadata)
|
||||
}
|
||||
|
||||
/// Compute average validation loss over the provided dataset.
|
||||
fn validate(&mut self, val_data: &[(Tensor, Tensor)]) -> Result<f64, MLError> {
|
||||
if val_data.is_empty() {
|
||||
return Err(MLError::ValidationError {
|
||||
message: "Empty validation dataset".to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
let mut total_loss = 0.0;
|
||||
let mut count = 0usize;
|
||||
|
||||
for (input, target) in val_data {
|
||||
let prediction = self.forward(input)?;
|
||||
let loss = self.compute_loss(&prediction, target)?;
|
||||
let loss_value = loss
|
||||
.to_dtype(candle_core::DType::F64)
|
||||
.and_then(|t| t.to_scalar::<f64>())
|
||||
.map_err(|e| MLError::ValidationError {
|
||||
message: format!("Failed to extract loss value: {}", e),
|
||||
})?;
|
||||
total_loss += loss_value;
|
||||
count += 1;
|
||||
}
|
||||
|
||||
let avg_loss = total_loss / count as f64;
|
||||
self.latest_metrics.val_loss = Some(avg_loss);
|
||||
|
||||
tracing::debug!("Liquid-CfC validation loss: {:.6}", avg_loss);
|
||||
|
||||
Ok(avg_loss)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::liquid::candle_cfc::DeviceConfig;
|
||||
use crate::training::unified_trainer::UnifiedTrainable;
|
||||
use candle_core::{DType, Tensor};
|
||||
|
||||
#[test]
|
||||
fn test_adapter_creation() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 8,
|
||||
hidden_size: 16,
|
||||
output_size: 3,
|
||||
backbone_hidden_sizes: vec![16],
|
||||
device: DeviceConfig::Cpu,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
assert_eq!(adapter.model_type(), "Liquid-CfC");
|
||||
assert_eq!(adapter.get_step(), 0);
|
||||
assert!(adapter.get_learning_rate() > 0.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_train_step() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 8,
|
||||
hidden_size: 16,
|
||||
output_size: 3,
|
||||
backbone_hidden_sizes: vec![16],
|
||||
device: DeviceConfig::Cpu,
|
||||
seq_len: 5,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
let device = adapter.device().clone();
|
||||
|
||||
let input = Tensor::randn(0f32, 1.0, (4, 5, 8), &device).unwrap();
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
assert_eq!(output.dims(), &[4, 3]);
|
||||
|
||||
let target = Tensor::zeros((4, 3), DType::F32, &device).unwrap();
|
||||
let loss = adapter.compute_loss(&output, &target).unwrap();
|
||||
let grad_norm = adapter.backward(&loss).unwrap();
|
||||
assert!(grad_norm >= 0.0);
|
||||
|
||||
adapter.optimizer_step().unwrap();
|
||||
assert_eq!(adapter.get_step(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_validate() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
seq_len: 3,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
let device = adapter.device().clone();
|
||||
|
||||
let val_data: Vec<(Tensor, Tensor)> = (0..3)
|
||||
.map(|_| {
|
||||
(
|
||||
Tensor::randn(0f32, 1.0, (2, 3, 4), &device).unwrap(),
|
||||
Tensor::zeros((2, 2), DType::F32, &device).unwrap(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
let val_loss = adapter.validate(&val_data).unwrap();
|
||||
assert!(val_loss.is_finite());
|
||||
assert!(val_loss >= 0.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_metrics() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
let metrics = adapter.collect_metrics();
|
||||
assert!(metrics.custom_metrics.contains_key("param_count"));
|
||||
assert!(metrics.custom_metrics.contains_key("num_parameters"));
|
||||
assert_eq!(metrics.learning_rate, 0.001);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_set_learning_rate() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
assert!((adapter.get_learning_rate() - 0.001).abs() < f64::EPSILON);
|
||||
|
||||
adapter.set_learning_rate(0.0001).unwrap();
|
||||
assert!((adapter.get_learning_rate() - 0.0001).abs() < f64::EPSILON);
|
||||
|
||||
// Invalid LR should error
|
||||
assert!(adapter.set_learning_rate(-0.1).is_err());
|
||||
assert!(adapter.set_learning_rate(0.0).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_zero_grad() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
seq_len: 3,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
let device = adapter.device().clone();
|
||||
|
||||
// Do a forward/backward to create grads
|
||||
let input = Tensor::randn(0f32, 1.0, (2, 3, 4), &device).unwrap();
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
let target = Tensor::zeros((2, 2), DType::F32, &device).unwrap();
|
||||
let loss = adapter.compute_loss(&output, &target).unwrap();
|
||||
adapter.backward(&loss).unwrap();
|
||||
|
||||
// zero_grad should clear stored grads
|
||||
adapter.zero_grad().unwrap();
|
||||
assert!(adapter.last_grads.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_checkpoint_roundtrip() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
seq_len: 3,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config.clone()).unwrap();
|
||||
let device = adapter.device().clone();
|
||||
|
||||
// Do a train step so step > 0
|
||||
let input = Tensor::randn(0f32, 1.0, (2, 3, 4), &device).unwrap();
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
let target = Tensor::zeros((2, 2), DType::F32, &device).unwrap();
|
||||
let loss = adapter.compute_loss(&output, &target).unwrap();
|
||||
adapter.backward(&loss).unwrap();
|
||||
adapter.optimizer_step().unwrap();
|
||||
assert_eq!(adapter.get_step(), 1);
|
||||
|
||||
// Save checkpoint
|
||||
let tmp_dir = std::env::temp_dir().join("liquid_cfc_test_ckpt");
|
||||
let _ = std::fs::create_dir_all(&tmp_dir);
|
||||
let ckpt_path = tmp_dir.join("test_ckpt");
|
||||
let ckpt_str = ckpt_path.to_str().unwrap();
|
||||
let saved = adapter.save_checkpoint(ckpt_str).unwrap();
|
||||
assert!(!saved.is_empty());
|
||||
|
||||
// Load into a fresh adapter
|
||||
let mut adapter2 = LiquidTrainableAdapter::new(config).unwrap();
|
||||
assert_eq!(adapter2.get_step(), 0);
|
||||
|
||||
let metadata = adapter2.load_checkpoint(ckpt_str).unwrap();
|
||||
assert_eq!(metadata.model_type, "Liquid-CfC");
|
||||
assert_eq!(metadata.step, 1);
|
||||
assert_eq!(adapter2.get_step(), 1);
|
||||
|
||||
// Cleanup
|
||||
let _ = std::fs::remove_file(format!("{}.safetensors", ckpt_str));
|
||||
let _ = std::fs::remove_file(format!("{}.json", ckpt_str));
|
||||
let _ = std::fs::remove_dir(&tmp_dir);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_validate_empty_errors() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
let result = adapter.validate(&[]);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_optimizer_step_without_backward_errors() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
let result = adapter.optimizer_step();
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_adapter_multiple_train_steps() {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: 4,
|
||||
hidden_size: 8,
|
||||
output_size: 2,
|
||||
backbone_hidden_sizes: vec![8],
|
||||
device: DeviceConfig::Cpu,
|
||||
seq_len: 3,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let mut adapter = LiquidTrainableAdapter::new(config).unwrap();
|
||||
let device = adapter.device().clone();
|
||||
|
||||
let input = Tensor::randn(0f32, 1.0, (4, 3, 4), &device).unwrap();
|
||||
let target = Tensor::zeros((4, 2), DType::F32, &device).unwrap();
|
||||
|
||||
let mut prev_loss = f64::MAX;
|
||||
for i in 0..5 {
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
let loss = adapter.compute_loss(&output, &target).unwrap();
|
||||
let loss_val: f64 = loss
|
||||
.to_dtype(DType::F64)
|
||||
.unwrap()
|
||||
.to_scalar()
|
||||
.unwrap();
|
||||
adapter.backward(&loss).unwrap();
|
||||
adapter.optimizer_step().unwrap();
|
||||
assert_eq!(adapter.get_step(), i + 1);
|
||||
|
||||
// Loss should generally decrease (not strictly, but should after a few steps)
|
||||
if i > 2 {
|
||||
// After several steps loss should be lower than the initial
|
||||
let _ = prev_loss; // avoid unused warning
|
||||
}
|
||||
prev_loss = loss_val;
|
||||
}
|
||||
// After 5 optimizer steps on the same batch, loss should have decreased
|
||||
// (not checking strictly since CfC dynamics make this non-trivial)
|
||||
assert!(prev_loss.is_finite());
|
||||
}
|
||||
}
|
||||
@@ -13,6 +13,7 @@ use crate::MLError;
|
||||
use common::trading::MarketRegime;
|
||||
|
||||
pub mod activation;
|
||||
pub mod adapter;
|
||||
pub mod candle_cfc;
|
||||
pub mod cells;
|
||||
pub mod network;
|
||||
@@ -24,6 +25,7 @@ mod tests;
|
||||
|
||||
// Re-export main types for external usage
|
||||
pub use activation::ActivationType;
|
||||
pub use adapter::LiquidTrainableAdapter;
|
||||
pub use candle_cfc::{CfCTrainConfig, DeviceConfig};
|
||||
pub use cells::{CfCConfig, LTCConfig};
|
||||
pub use network::{LayerConfig, LiquidNetwork, LiquidNetworkConfig, OutputLayerConfig};
|
||||
|
||||
Reference in New Issue
Block a user