diff --git a/ml/src/tgnn/mod.rs b/ml/src/tgnn/mod.rs index 4b205ba88..627edefdb 100644 --- a/ml/src/tgnn/mod.rs +++ b/ml/src/tgnn/mod.rs @@ -21,6 +21,7 @@ pub mod gating; pub mod graph; pub mod message_passing; +pub mod trainable_adapter; pub mod traits; pub mod types; diff --git a/ml/src/tgnn/trainable_adapter.rs b/ml/src/tgnn/trainable_adapter.rs new file mode 100644 index 000000000..5ac373762 --- /dev/null +++ b/ml/src/tgnn/trainable_adapter.rs @@ -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, + /// 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, + /// 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 { + 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 { + // 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 { + // 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 { + 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::()) + .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 = self.loss_history.iter().rev().take(100).copied().collect(); + metrics.loss = recent.iter().sum::() / 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 { + 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 = 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 { + 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 { + 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::().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::().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"); + } +}