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>
642 lines
22 KiB
Rust
642 lines
22 KiB
Rust
//! 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());
|
|
}
|
|
}
|