Merge branch 'feat/tggn-fullstack' — TGGN UnifiedTrainable adapter

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-23 01:01:03 +01:00
2 changed files with 598 additions and 0 deletions

View File

@@ -21,6 +21,7 @@
pub mod gating;
pub mod graph;
pub mod message_passing;
pub mod trainable_adapter;
pub mod traits;
pub mod types;

View 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");
}
}