Merge branch 'feat/tggn-fullstack' — TGGN UnifiedTrainable adapter
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -21,6 +21,7 @@
|
||||
pub mod gating;
|
||||
pub mod graph;
|
||||
pub mod message_passing;
|
||||
pub mod trainable_adapter;
|
||||
pub mod traits;
|
||||
pub mod types;
|
||||
|
||||
|
||||
597
ml/src/tgnn/trainable_adapter.rs
Normal file
597
ml/src/tgnn/trainable_adapter.rs
Normal file
@@ -0,0 +1,597 @@
|
||||
//! UnifiedTrainable adapter for TGGN (Temporal Graph Gated Network)
|
||||
//!
|
||||
//! Wraps a candle-based projection network to provide the UnifiedTrainable
|
||||
//! interface for TGGN. The projection network maps flattened graph features
|
||||
//! through hidden layers, enabling gradient-based training via the unified
|
||||
//! training orchestrator.
|
||||
//!
|
||||
//! Architecture: input_linear(node_dim -> hidden_dim) -> ReLU -> output_linear(hidden_dim -> 1)
|
||||
|
||||
use candle_core::{backprop::GradStore, DType, Device, Module, Tensor};
|
||||
use candle_nn::{linear, AdamW, Linear, Optimizer, ParamsAdamW, VarBuilder, VarMap};
|
||||
use std::collections::HashMap;
|
||||
|
||||
use super::TGGNConfig;
|
||||
use crate::training::unified_trainer::{checkpoint, CheckpointMetadata, TrainingMetrics, UnifiedTrainable};
|
||||
use crate::MLError;
|
||||
|
||||
/// Adapter wrapping TGGN with a candle-based projection network for unified training.
|
||||
///
|
||||
/// The projection network has two linear layers:
|
||||
/// - `input_linear`: projects from `node_dim` to `hidden_dim`
|
||||
/// - `output_linear`: projects from `hidden_dim` to 1 (scalar prediction)
|
||||
///
|
||||
/// Training uses AdamW with gradient tracking via `GradStore`.
|
||||
pub struct TGGNTrainableAdapter {
|
||||
/// TGGN configuration
|
||||
config: TGGNConfig,
|
||||
/// Candle variable map holding learnable parameters
|
||||
var_map: VarMap,
|
||||
/// Input projection layer (node_dim -> hidden_dim)
|
||||
input_linear: Linear,
|
||||
/// Output projection layer (hidden_dim -> 1)
|
||||
output_linear: Linear,
|
||||
/// AdamW optimizer
|
||||
optimizer: AdamW,
|
||||
/// Gradient store from last backward pass (consumed by optimizer_step)
|
||||
grads: Option<GradStore>,
|
||||
/// Device (CPU or CUDA)
|
||||
device: Device,
|
||||
/// Current learning rate
|
||||
learning_rate: f64,
|
||||
/// Current training step
|
||||
step: usize,
|
||||
/// Latest training metrics
|
||||
latest_metrics: TrainingMetrics,
|
||||
/// Loss history for rolling average
|
||||
loss_history: Vec<f64>,
|
||||
/// Last computed gradient norm
|
||||
last_grad_norm: f64,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for TGGNTrainableAdapter {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("TGGNTrainableAdapter")
|
||||
.field("config", &self.config)
|
||||
.field("device", &format!("{:?}", self.device))
|
||||
.field("learning_rate", &self.learning_rate)
|
||||
.field("step", &self.step)
|
||||
.field("loss_history_len", &self.loss_history.len())
|
||||
.field("last_grad_norm", &self.last_grad_norm)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl TGGNTrainableAdapter {
|
||||
/// Create a new TGGN trainable adapter with projection network.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `config` - TGGN configuration specifying dimensions
|
||||
/// * `device` - Device to create tensors on (CPU or CUDA)
|
||||
///
|
||||
/// # Returns
|
||||
/// Initialized adapter ready for training
|
||||
pub fn new(config: TGGNConfig, device: &Device) -> Result<Self, MLError> {
|
||||
let var_map = VarMap::new();
|
||||
let vb = VarBuilder::from_varmap(&var_map, DType::F32, 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)))?;
|
||||
|
||||
let output_linear = linear(config.hidden_dim, 1, vb.pp("output"))
|
||||
.map_err(|e| MLError::ModelError(format!("Failed to create output linear: {}", e)))?;
|
||||
|
||||
let learning_rate = 1e-3;
|
||||
let all_vars = var_map.all_vars();
|
||||
let optimizer = AdamW::new(
|
||||
all_vars,
|
||||
ParamsAdamW {
|
||||
lr: learning_rate,
|
||||
beta1: 0.9,
|
||||
beta2: 0.999,
|
||||
eps: 1e-8,
|
||||
weight_decay: 1e-4,
|
||||
},
|
||||
)
|
||||
.map_err(|e| MLError::ModelError(format!("Failed to create AdamW optimizer: {}", e)))?;
|
||||
|
||||
Ok(Self {
|
||||
config,
|
||||
var_map,
|
||||
input_linear,
|
||||
output_linear,
|
||||
optimizer,
|
||||
grads: None,
|
||||
device: device.clone(),
|
||||
learning_rate,
|
||||
step: 0,
|
||||
latest_metrics: TrainingMetrics::default(),
|
||||
loss_history: Vec::new(),
|
||||
last_grad_norm: 0.0,
|
||||
})
|
||||
}
|
||||
|
||||
/// Access the underlying TGGN configuration.
|
||||
pub fn tggn_config(&self) -> &TGGNConfig {
|
||||
&self.config
|
||||
}
|
||||
}
|
||||
|
||||
impl UnifiedTrainable for TGGNTrainableAdapter {
|
||||
fn model_type(&self) -> &str {
|
||||
"TGGN"
|
||||
}
|
||||
|
||||
fn device(&self) -> &Device {
|
||||
&self.device
|
||||
}
|
||||
|
||||
fn forward(&mut self, input: &Tensor) -> Result<Tensor, MLError> {
|
||||
// input: [batch, node_dim]
|
||||
// input_linear: node_dim -> hidden_dim
|
||||
let hidden = self.input_linear.forward(input).map_err(|e| {
|
||||
MLError::ModelError(format!("Input linear forward failed: {}", e))
|
||||
})?;
|
||||
|
||||
// ReLU activation
|
||||
let activated = hidden.relu().map_err(|e| {
|
||||
MLError::ModelError(format!("ReLU activation failed: {}", e))
|
||||
})?;
|
||||
|
||||
// output_linear: hidden_dim -> 1
|
||||
let output = self.output_linear.forward(&activated).map_err(|e| {
|
||||
MLError::ModelError(format!("Output linear forward failed: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn compute_loss(&self, predictions: &Tensor, targets: &Tensor) -> Result<Tensor, MLError> {
|
||||
// MSE loss: mean((predictions - targets)^2)
|
||||
let diff = predictions.sub(targets).map_err(|e| {
|
||||
MLError::ModelError(format!("Loss subtraction failed: {}", e))
|
||||
})?;
|
||||
let squared = diff.powf(2.0).map_err(|e| {
|
||||
MLError::ModelError(format!("Loss squaring failed: {}", e))
|
||||
})?;
|
||||
let loss = squared.mean_all().map_err(|e| {
|
||||
MLError::ModelError(format!("Loss mean failed: {}", e))
|
||||
})?;
|
||||
Ok(loss)
|
||||
}
|
||||
|
||||
fn backward(&mut self, loss: &Tensor) -> Result<f64, MLError> {
|
||||
let grads = loss.backward().map_err(|e| {
|
||||
MLError::TrainingError(format!("Backward pass failed: {}", e))
|
||||
})?;
|
||||
|
||||
// Compute gradient norm across all variables for monitoring
|
||||
let mut grad_norm_sq = 0.0;
|
||||
let vars_lock = self
|
||||
.var_map
|
||||
.data()
|
||||
.lock()
|
||||
.map_err(|e| MLError::LockError(format!("Failed to lock var_map: {}", e)))?;
|
||||
|
||||
for (_name, var) in vars_lock.iter() {
|
||||
if let Some(grad) = grads.get(var.as_tensor()) {
|
||||
let norm = grad
|
||||
.sqr()
|
||||
.and_then(|s| s.sum_all())
|
||||
.and_then(|s| s.to_scalar::<f32>())
|
||||
.map_err(|e| {
|
||||
MLError::ModelError(format!("Failed to compute grad norm: {}", e))
|
||||
})?;
|
||||
grad_norm_sq += norm as f64;
|
||||
}
|
||||
}
|
||||
drop(vars_lock);
|
||||
|
||||
let grad_norm = grad_norm_sq.sqrt();
|
||||
self.last_grad_norm = grad_norm;
|
||||
self.latest_metrics.grad_norm = Some(grad_norm);
|
||||
self.grads = Some(grads);
|
||||
|
||||
Ok(grad_norm)
|
||||
}
|
||||
|
||||
fn optimizer_step(&mut self) -> Result<(), MLError> {
|
||||
if let Some(grads) = self.grads.take() {
|
||||
self.optimizer.step(&grads).map_err(|e| {
|
||||
MLError::TrainingError(format!("Optimizer step failed: {}", e))
|
||||
})?;
|
||||
}
|
||||
self.step += 1;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn zero_grad(&mut self) -> Result<(), MLError> {
|
||||
// Gradients are implicitly zeroed by creating a new GradStore in backward().
|
||||
// No explicit zeroing needed with candle's approach.
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn get_learning_rate(&self) -> f64 {
|
||||
self.learning_rate
|
||||
}
|
||||
|
||||
fn set_learning_rate(&mut self, lr: f64) -> Result<(), MLError> {
|
||||
self.learning_rate = lr;
|
||||
// Recreate optimizer with new learning rate
|
||||
let all_vars = self.var_map.all_vars();
|
||||
self.optimizer = AdamW::new(
|
||||
all_vars,
|
||||
ParamsAdamW {
|
||||
lr,
|
||||
beta1: 0.9,
|
||||
beta2: 0.999,
|
||||
eps: 1e-8,
|
||||
weight_decay: 1e-4,
|
||||
},
|
||||
)
|
||||
.map_err(|e| {
|
||||
MLError::ModelError(format!("Failed to recreate optimizer with new lr: {}", e))
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn get_step(&self) -> usize {
|
||||
self.step
|
||||
}
|
||||
|
||||
fn collect_metrics(&self) -> TrainingMetrics {
|
||||
let mut metrics = self.latest_metrics.clone();
|
||||
|
||||
// Rolling average of recent losses
|
||||
if !self.loss_history.is_empty() {
|
||||
let recent: Vec<f64> = self.loss_history.iter().rev().take(100).copied().collect();
|
||||
metrics.loss = recent.iter().sum::<f64>() / recent.len() as f64;
|
||||
}
|
||||
|
||||
metrics.learning_rate = self.learning_rate;
|
||||
|
||||
// TGGN-specific custom metrics
|
||||
metrics
|
||||
.custom_metrics
|
||||
.insert("training_steps".to_string(), self.step as f64);
|
||||
metrics.custom_metrics.insert(
|
||||
"loss_history_len".to_string(),
|
||||
self.loss_history.len() as f64,
|
||||
);
|
||||
metrics.custom_metrics.insert(
|
||||
"node_dim".to_string(),
|
||||
self.config.node_dim as f64,
|
||||
);
|
||||
metrics.custom_metrics.insert(
|
||||
"hidden_dim".to_string(),
|
||||
self.config.hidden_dim as f64,
|
||||
);
|
||||
metrics
|
||||
.custom_metrics
|
||||
.insert("num_layers".to_string(), self.config.num_layers as f64);
|
||||
|
||||
metrics
|
||||
}
|
||||
|
||||
fn save_checkpoint(&self, checkpoint_path: &str) -> Result<String, MLError> {
|
||||
let metadata = CheckpointMetadata {
|
||||
model_type: "TGGN".to_string(),
|
||||
version: "1.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 serialize config: {}", e),
|
||||
}
|
||||
})?,
|
||||
metrics: self.collect_metrics(),
|
||||
};
|
||||
|
||||
// Save JSON metadata
|
||||
checkpoint::save_metadata(&metadata, checkpoint_path)?;
|
||||
|
||||
// Save model weights via safetensors
|
||||
let safetensors_path = format!("{}.safetensors", checkpoint_path);
|
||||
let vars_lock = self
|
||||
.var_map
|
||||
.data()
|
||||
.lock()
|
||||
.map_err(|e| MLError::LockError(format!("Failed to lock var_map: {}", e)))?;
|
||||
|
||||
let mut tensors: HashMap<String, Tensor> = HashMap::new();
|
||||
for (name, var) in vars_lock.iter() {
|
||||
tensors.insert(name.clone(), var.as_tensor().clone());
|
||||
}
|
||||
drop(vars_lock);
|
||||
|
||||
candle_core::safetensors::save(&tensors, &safetensors_path).map_err(|e| {
|
||||
MLError::CheckpointError(format!("Failed to save safetensors: {}", e))
|
||||
})?;
|
||||
|
||||
tracing::info!(
|
||||
"Saved TGGN checkpoint to {} (step {})",
|
||||
checkpoint_path,
|
||||
self.step
|
||||
);
|
||||
|
||||
Ok(checkpoint_path.to_string())
|
||||
}
|
||||
|
||||
fn load_checkpoint(&mut self, checkpoint_path: &str) -> Result<CheckpointMetadata, MLError> {
|
||||
let metadata = checkpoint::load_metadata(checkpoint_path)?;
|
||||
|
||||
if metadata.model_type != "TGGN" {
|
||||
return Err(MLError::CheckpointError(format!(
|
||||
"Invalid model type in checkpoint: expected TGGN, got {}",
|
||||
metadata.model_type
|
||||
)));
|
||||
}
|
||||
|
||||
// Load weights from safetensors
|
||||
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)),
|
||||
)?;
|
||||
|
||||
// Set loaded tensors into VarMap
|
||||
let vars_lock = self
|
||||
.var_map
|
||||
.data()
|
||||
.lock()
|
||||
.map_err(|e| MLError::LockError(format!("Failed to lock var_map: {}", e)))?;
|
||||
|
||||
for (name, tensor) in &tensors {
|
||||
if let Some(var) = vars_lock.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);
|
||||
}
|
||||
}
|
||||
drop(vars_lock);
|
||||
|
||||
// Restore training state
|
||||
self.step = metadata.step;
|
||||
self.latest_metrics = metadata.metrics.clone();
|
||||
|
||||
tracing::info!(
|
||||
"Loaded TGGN checkpoint from {} (step {})",
|
||||
checkpoint_path,
|
||||
metadata.step
|
||||
);
|
||||
|
||||
Ok(metadata)
|
||||
}
|
||||
|
||||
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_val = loss.to_scalar::<f32>().map_err(|e| MLError::ValidationError {
|
||||
message: format!("Failed to extract loss value: {}", e),
|
||||
})?;
|
||||
total_loss += loss_val as f64;
|
||||
count += 1;
|
||||
}
|
||||
|
||||
let avg_loss = if count > 0 {
|
||||
total_loss / count as f64
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
self.latest_metrics.val_loss = Some(avg_loss);
|
||||
|
||||
tracing::debug!("TGGN validation loss: {:.6}", avg_loss);
|
||||
|
||||
Ok(avg_loss)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn make_config() -> TGGNConfig {
|
||||
TGGNConfig {
|
||||
max_nodes: 16,
|
||||
max_edges: 32,
|
||||
node_dim: 8,
|
||||
edge_dim: 4,
|
||||
hidden_dim: 16,
|
||||
num_layers: 2,
|
||||
temporal_decay: 0.99,
|
||||
update_frequency_ns: 1_000_000,
|
||||
use_simd: false,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_model_type() {
|
||||
let cfg = make_config();
|
||||
let adapter = TGGNTrainableAdapter::new(cfg, &Device::Cpu).unwrap();
|
||||
assert_eq!(adapter.model_type(), "TGGN");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_device() {
|
||||
let cfg = make_config();
|
||||
let adapter = TGGNTrainableAdapter::new(cfg, &Device::Cpu).unwrap();
|
||||
assert!(matches!(adapter.device(), &Device::Cpu));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_forward_shape() {
|
||||
let cfg = make_config();
|
||||
let mut adapter = TGGNTrainableAdapter::new(cfg.clone(), &Device::Cpu).unwrap();
|
||||
|
||||
// batch=2, node_dim=8
|
||||
let input = Tensor::zeros(&[2, cfg.node_dim], DType::F32, &Device::Cpu).unwrap();
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
|
||||
let dims = output.shape().dims();
|
||||
assert_eq!(dims.len(), 2);
|
||||
assert_eq!(dims[0], 2); // batch
|
||||
assert_eq!(dims[1], 1); // scalar prediction
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_compute_loss() {
|
||||
let cfg = make_config();
|
||||
let adapter = TGGNTrainableAdapter::new(cfg, &Device::Cpu).unwrap();
|
||||
|
||||
let preds = Tensor::new(&[[1.0f32], [2.0]], &Device::Cpu).unwrap();
|
||||
let targets = Tensor::new(&[[1.5f32], [2.5]], &Device::Cpu).unwrap();
|
||||
|
||||
let loss = adapter.compute_loss(&preds, &targets).unwrap();
|
||||
let loss_val: f32 = loss.to_scalar().unwrap();
|
||||
|
||||
// MSE of (0.5^2 + 0.5^2)/2 = 0.25
|
||||
assert!((loss_val - 0.25).abs() < 1e-5, "Expected ~0.25, got {}", loss_val);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_backward_returns_grad_norm() {
|
||||
let cfg = make_config();
|
||||
let mut adapter = TGGNTrainableAdapter::new(cfg.clone(), &Device::Cpu).unwrap();
|
||||
|
||||
let input = Tensor::randn(0.0f32, 1.0, &[2, cfg.node_dim], &Device::Cpu).unwrap();
|
||||
let targets = Tensor::randn(0.0f32, 1.0, &[2, 1], &Device::Cpu).unwrap();
|
||||
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
let loss = adapter.compute_loss(&output, &targets).unwrap();
|
||||
let grad_norm = adapter.backward(&loss).unwrap();
|
||||
|
||||
assert!(grad_norm >= 0.0, "Gradient norm should be non-negative");
|
||||
assert!(adapter.grads.is_some(), "Grads should be stored after backward");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_train_step_cycle() {
|
||||
let cfg = make_config();
|
||||
let mut adapter = TGGNTrainableAdapter::new(cfg.clone(), &Device::Cpu).unwrap();
|
||||
|
||||
assert_eq!(adapter.get_step(), 0);
|
||||
|
||||
// Full train cycle: forward -> loss -> backward -> optimizer_step
|
||||
let input = Tensor::randn(0.0f32, 1.0, &[4, cfg.node_dim], &Device::Cpu).unwrap();
|
||||
let targets = Tensor::randn(0.0f32, 1.0, &[4, 1], &Device::Cpu).unwrap();
|
||||
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
let loss = adapter.compute_loss(&output, &targets).unwrap();
|
||||
let loss_val: f64 = loss.to_scalar::<f32>().unwrap() as f64;
|
||||
|
||||
adapter.backward(&loss).unwrap();
|
||||
adapter.optimizer_step().unwrap();
|
||||
|
||||
assert_eq!(adapter.get_step(), 1);
|
||||
|
||||
// Record loss and check metrics
|
||||
adapter.loss_history.push(loss_val);
|
||||
let metrics = adapter.collect_metrics();
|
||||
assert!(metrics.loss > 0.0 || metrics.loss == 0.0);
|
||||
assert_eq!(
|
||||
metrics.custom_metrics.get("training_steps").copied(),
|
||||
Some(1.0)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_learning_rate_get_set() {
|
||||
let cfg = make_config();
|
||||
let mut adapter = TGGNTrainableAdapter::new(cfg, &Device::Cpu).unwrap();
|
||||
|
||||
let original_lr = adapter.get_learning_rate();
|
||||
assert!((original_lr - 1e-3).abs() < 1e-10);
|
||||
|
||||
adapter.set_learning_rate(5e-4).unwrap();
|
||||
assert!((adapter.get_learning_rate() - 5e-4).abs() < 1e-10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_collect_metrics() {
|
||||
let cfg = make_config();
|
||||
let adapter = TGGNTrainableAdapter::new(cfg, &Device::Cpu).unwrap();
|
||||
|
||||
let metrics = adapter.collect_metrics();
|
||||
assert!(metrics.custom_metrics.contains_key("training_steps"));
|
||||
assert!(metrics.custom_metrics.contains_key("node_dim"));
|
||||
assert!(metrics.custom_metrics.contains_key("hidden_dim"));
|
||||
assert!(metrics.custom_metrics.contains_key("num_layers"));
|
||||
assert_eq!(metrics.custom_metrics.get("node_dim").copied(), Some(8.0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_checkpoint_roundtrip() {
|
||||
let cfg = make_config();
|
||||
let mut adapter = TGGNTrainableAdapter::new(cfg.clone(), &Device::Cpu).unwrap();
|
||||
|
||||
// Run a training step to have non-zero state
|
||||
let input = Tensor::randn(0.0f32, 1.0, &[2, cfg.node_dim], &Device::Cpu).unwrap();
|
||||
let targets = Tensor::randn(0.0f32, 1.0, &[2, 1], &Device::Cpu).unwrap();
|
||||
let output = adapter.forward(&input).unwrap();
|
||||
let loss = adapter.compute_loss(&output, &targets).unwrap();
|
||||
adapter.backward(&loss).unwrap();
|
||||
adapter.optimizer_step().unwrap();
|
||||
|
||||
// Save checkpoint
|
||||
let tmp_dir = std::env::temp_dir();
|
||||
let checkpoint_path = tmp_dir.join("tggn_test_ckpt");
|
||||
let path_str = checkpoint_path.to_str().unwrap();
|
||||
adapter.save_checkpoint(path_str).unwrap();
|
||||
|
||||
// Load into fresh adapter
|
||||
let mut adapter2 = TGGNTrainableAdapter::new(cfg, &Device::Cpu).unwrap();
|
||||
let metadata = adapter2.load_checkpoint(path_str).unwrap();
|
||||
|
||||
assert_eq!(metadata.model_type, "TGGN");
|
||||
assert_eq!(metadata.step, 1);
|
||||
assert_eq!(adapter2.get_step(), 1);
|
||||
|
||||
// Clean up checkpoint files
|
||||
let _ = std::fs::remove_file(format!("{}.json", path_str));
|
||||
let _ = std::fs::remove_file(format!("{}.safetensors", path_str));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate() {
|
||||
let cfg = make_config();
|
||||
let mut adapter = TGGNTrainableAdapter::new(cfg.clone(), &Device::Cpu).unwrap();
|
||||
|
||||
let val_data: Vec<(Tensor, Tensor)> = (0..3)
|
||||
.map(|_| {
|
||||
let input =
|
||||
Tensor::randn(0.0f32, 1.0, &[2, cfg.node_dim], &Device::Cpu).unwrap();
|
||||
let target = Tensor::randn(0.0f32, 1.0, &[2, 1], &Device::Cpu).unwrap();
|
||||
(input, target)
|
||||
})
|
||||
.collect();
|
||||
|
||||
let val_loss = adapter.validate(&val_data).unwrap();
|
||||
assert!(val_loss >= 0.0, "Validation loss should be non-negative");
|
||||
assert!(
|
||||
adapter.latest_metrics.val_loss.is_some(),
|
||||
"val_loss should be set after validate"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_empty_errors() {
|
||||
let cfg = make_config();
|
||||
let mut adapter = TGGNTrainableAdapter::new(cfg, &Device::Cpu).unwrap();
|
||||
|
||||
let result = adapter.validate(&[]);
|
||||
assert!(result.is_err(), "Validating empty data should error");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user