use candle_core::Var;
use error::KohoError;
use math::tensors::Matrix;
use nn::{
diffuse::DiffusionLayer,
loss::{LossFn, LossKind},
metrics::{EpochMetrics, TrainingMetrics},
optim::{OptimKind, OptimizerParams},
};
use crate::{math::sheaf::CellularSheaf, nn::optim::create_optimizer};
pub mod error;
pub mod math;
pub mod nn;
pub trait Parameterized {
fn parameters(&self) -> Vec<Var>;
fn parameters_mut(&mut self) -> Vec<&mut Var>;
}
pub struct SheafNN {
sheaf: CellularSheaf,
layers: Vec<DiffusionLayer>,
loss_fn: LossFn,
k: usize,
down_included: bool,
}
impl SheafNN {
pub fn init(
k: usize,
down_laplacian_included: bool,
loss: LossKind,
sheaf: CellularSheaf,
) -> Self {
let loss_fn = LossFn::new(loss);
Self {
sheaf,
layers: Vec::new(),
loss_fn,
k,
down_included: down_laplacian_included,
}
}
pub fn sequential(&mut self, layers: Vec<DiffusionLayer>) {
self.layers = layers;
}
pub fn forward(&self, input: Matrix, down_included: bool) -> Result<Matrix, KohoError> {
let mut output = input;
for i in &self.layers {
output = i.diffuse(&self.sheaf, self.k, output, down_included)?;
}
Ok(output)
}
pub fn train(
&mut self,
data: &[(Matrix, Matrix)],
epochs: usize,
down_included: bool,
optimizer_kind: OptimKind,
lr: f64,
optimizer_params: OptimizerParams,
) -> Result<TrainingMetrics, KohoError> {
let mut optimizer =
create_optimizer(optimizer_kind, self.parameters_mut(), lr, optimizer_params)?;
let mut metrics = TrainingMetrics::new(epochs);
for epoch in 1..=epochs {
let mut total_loss = 0.0_f32;
for (input, target) in data {
let output = self.forward(input.clone(), down_included)?;
let loss_tensor = self.loss_fn.compute(output.inner(), target.inner())?;
let loss_val = loss_tensor.to_scalar::<f32>().unwrap_or(f32::NAN);
total_loss += loss_val;
let grads = loss_tensor.backward()?;
optimizer.step(&grads, self.parameters_mut())?;
}
let avg_loss = total_loss / (data.len() as f32);
metrics.push(EpochMetrics::new(epoch, avg_loss));
}
Ok(metrics)
}
pub fn train_debug(
&mut self,
data: &[(Matrix, Matrix)],
epochs: usize,
down_included: bool,
optimizer_kind: OptimKind,
lr: f64,
optimizer_params: OptimizerParams,
) -> Result<TrainingMetrics, KohoError> {
println!("=== Training Debug Info ===");
let params = self.parameters();
println!("Total parameters: {}", params.len());
for (i, param) in params.iter().enumerate() {
let param_data = param.as_tensor().flatten_all()?;
let param_vec = param_data.to_vec1::<f32>()?;
println!(
"Parameter {i}: shape={:?}, first_few_values={:?}",
param.shape(),
¶m_vec[..param_vec.len().min(5)]
);
}
let mut optimizer =
create_optimizer(optimizer_kind, self.parameters_mut(), lr, optimizer_params)?;
let mut metrics = TrainingMetrics::new(epochs);
for epoch in 1..=epochs {
let mut total_loss = 0.0_f32;
for (batch_idx, (input, target)) in data.iter().enumerate() {
println!("\nEpoch {epoch}, Batch {batch_idx}");
let input_data = input.inner().flatten_all()?;
let target_data = target.inner().flatten_all()?;
let input_vec = input_data.to_vec1::<f32>()?;
let target_vec = target_data.to_vec1::<f32>()?;
println!("Input: {input_vec:?}");
println!("Target: {target_vec:?}");
let output = self.forward(input.clone(), down_included)?;
let output_data = output.inner().flatten_all()?;
let output_vec = output_data.to_vec1::<f32>()?;
println!("Output: {output_vec:?}");
let loss_tensor = self.loss_fn.compute(output.inner(), target.inner())?;
let loss_val = loss_tensor.to_scalar::<f32>().unwrap_or(f32::NAN);
total_loss += loss_val;
println!("Loss: {loss_val}");
println!("Loss tensor shape: {:?}", loss_tensor.shape());
println!("Loss tensor dtype: {:?}", loss_tensor.dtype());
println!("Computing gradients...");
let grads = loss_tensor.backward()?;
let params_mut = self.parameters_mut();
println!("Checking gradients for {} parameters:", params_mut.len());
for (i, param) in params_mut.iter().enumerate() {
if let Some(grad) = grads.get(param) {
let grad_data = grad.flatten_all()?;
let grad_vec = grad_data.to_vec1::<f32>()?;
let grad_norm = grad_vec.iter().map(|x| x * x).sum::<f32>().sqrt();
println!(
" Param {i}: grad_norm={grad_norm}, first_few_grads={:?}",
&grad_vec[..grad_vec.len().min(3)]
);
} else {
println!(" Param {i}: NO GRADIENT FOUND");
}
}
println!("Applying optimizer step...");
let params_before: Vec<_> = self
.parameters_mut()
.iter()
.map(|p| {
p.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
optimizer.step(&grads, self.parameters_mut())?;
let params_after: Vec<_> = self
.parameters_mut()
.iter()
.map(|p| {
p.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
for (i, (before, after)) in
params_before.iter().zip(params_after.iter()).enumerate()
{
let diff_norm: f32 = before
.iter()
.zip(after.iter())
.map(|(b, a)| (b - a).powi(2))
.sum::<f32>()
.sqrt();
println!(" Param {i} change norm: {diff_norm}");
}
if epoch <= 3 {
println!("--- End batch {batch_idx} ---");
}
}
let avg_loss = total_loss / (data.len() as f32);
metrics.push(EpochMetrics::new(epoch, avg_loss));
if epoch <= 10 || epoch % 10 == 0 {
println!("Epoch {epoch}: avg_loss = {avg_loss}");
}
}
Ok(metrics)
}
}
impl Parameterized for SheafNN {
fn parameters(&self) -> Vec<Var> {
let mut out = self
.layers
.iter()
.flat_map(|layer| layer.parameters())
.collect::<Vec<_>>();
if self.sheaf.learned {
out.extend(self.sheaf.parameters(self.k, self.down_included));
}
out
}
fn parameters_mut(&mut self) -> Vec<&mut Var> {
let learned = self.sheaf.learned;
let mut out = self
.layers
.iter_mut()
.flat_map(|layer| layer.parameters_mut())
.collect::<Vec<_>>();
if learned {
out.extend(self.sheaf.parameters_mut(self.k, self.down_included));
}
out
}
}
#[cfg(test)]
mod integration_tests {
use super::*;
use crate::{
math::{
cell::Cell,
sheaf::{CellularSheaf, Section},
tensors::Matrix,
},
nn::{
activate::Activations,
diffuse::DiffusionLayer,
loss::LossKind,
optim::{OptimKind, OptimizerParams},
},
};
use candle_core::{DType, Device};
fn create_triangle_sheaf() -> Result<CellularSheaf, KohoError> {
let mut sheaf = CellularSheaf::init(DType::F32, Device::Cpu, true);
let v0_data = Section::new(&[1.0f32], 1, Device::Cpu, DType::F32)?;
let (_, v0_idx) = sheaf.attach(Cell::new(0), v0_data, None, None)?;
let v1_data = Section::new(&[0.0f32], 1, Device::Cpu, DType::F32)?;
let (_, v1_idx) = sheaf.attach(Cell::new(0), v1_data, None, None)?;
let v2_data = Section::new(&[0.0f32], 1, Device::Cpu, DType::F32)?;
let (_, v2_idx) = sheaf.attach(Cell::new(0), v2_data, None, None)?;
let e0_data = Section::new(&[0.5f32], 1, Device::Cpu, DType::F32)?;
let (_, e0_idx) = sheaf.attach(Cell::new(1), e0_data, None, Some(&[v0_idx, v1_idx]))?;
let e1_data = Section::new(&[0.5f32], 1, Device::Cpu, DType::F32)?;
let (_, e1_idx) = sheaf.attach(Cell::new(1), e1_data, None, Some(&[v1_idx, v2_idx]))?;
let e2_data = Section::new(&[0.5f32], 1, Device::Cpu, DType::F32)?;
let (_, e2_idx) = sheaf.attach(Cell::new(1), e2_data, None, Some(&[v2_idx, v0_idx]))?;
let f0_data = Section::new(&[0.0f32], 1, Device::Cpu, DType::F32)?;
let (_, _f0_idx) =
sheaf.attach(Cell::new(2), f0_data, None, Some(&[e0_idx, e1_idx, e2_idx]))?;
sheaf.generate_initial_restrictions(0.1)?;
println!("uppers: {:?}", sheaf.cells.cells[0][0].upper);
Ok(sheaf)
}
#[test]
fn test_triangle_diffusion_learning() -> Result<(), KohoError> {
let sheaf = create_triangle_sheaf()?;
let input = sheaf.get_k_cochain(0)?;
let target_data = vec![0.8f32, 0.6f32, 0.4f32];
let target = Matrix::from_slice(&target_data, 1, 3, Device::Cpu, DType::F32)?;
let training_data = vec![(input, target)];
let mut network = SheafNN::init(0, false, LossKind::MSE, sheaf);
let diffusion_layer = DiffusionLayer::new(0, Activations::Linear, &network.sheaf)?;
network.sequential(vec![diffusion_layer]);
let metrics = network.train_debug(
&training_data,
100,
false,
OptimKind::Adam,
0.01,
OptimizerParams::Else,
)?;
assert!(
metrics.final_loss < metrics.epochs[0].loss,
"Training should reduce loss over time"
);
let initial_input = network.sheaf.get_k_cochain(0)?;
let output = network.forward(initial_input, false)?;
let input_vals = network.sheaf.get_k_cochain(0)?.inner().to_vec2::<f32>()?;
let output_vals = output.inner().to_vec2::<f32>()?;
println!("Input: {input_vals:?}");
println!("Output: {output_vals:?}");
println!("Final loss: {}", metrics.final_loss);
assert_eq!(output_vals.len(), 1, "Output should have 1 feature");
assert_eq!(
output_vals[0].len(),
3,
"Each vertex should have 3 vertices"
);
Ok(())
}
#[test]
fn test_edge_diffusion_learning() -> Result<(), KohoError> {
let sheaf = create_triangle_sheaf()?;
let input = sheaf.get_k_cochain(1)?;
println!("got edges");
let target_data = vec![0.5f32, 0.3f32, 0.7f32];
let target = Matrix::from_slice(&target_data, 1, 3, Device::Cpu, DType::F32)?;
let training_data = vec![(input, target)];
let mut network = SheafNN::init(1, true, LossKind::MSE, sheaf); let diffusion_layer = DiffusionLayer::new(1, Activations::Tanh, &network.sheaf)?;
network.sequential(vec![diffusion_layer]);
let metrics = network.train_debug(
&training_data,
200,
true,
OptimKind::Adam,
0.2,
OptimizerParams::Else,
)?;
println!("Edge diffusion final loss: {}", metrics.final_loss);
assert!(metrics.final_loss < 1.0, "Loss should be reasonable");
Ok(())
}
#[test]
fn test_learned_vs_fixed_restrictions() -> Result<(), KohoError> {
let sheaf_learned = create_triangle_sheaf()?;
assert!(sheaf_learned.learned, "Sheaf should have learned=true");
let mut sheaf_fixed = create_triangle_sheaf()?;
sheaf_fixed.learned = false;
let input = sheaf_learned.get_k_cochain(0)?;
let target_data = vec![0.8f32, 0.6f32, 0.4f32];
let target = Matrix::from_slice(&target_data, 1, 3, Device::Cpu, DType::F32)?;
let training_data = vec![(input.clone(), target.clone())];
let mut network_learned = SheafNN::init(0, false, LossKind::MSE, sheaf_learned);
let layer_learned = DiffusionLayer::new(0, Activations::Linear, &network_learned.sheaf)?;
network_learned.sequential(vec![layer_learned]);
let metrics_learned = network_learned.train(
&training_data,
50,
false,
OptimKind::Adam,
0.01,
OptimizerParams::Else,
)?;
let mut network_fixed = SheafNN::init(0, false, LossKind::MSE, sheaf_fixed);
let layer_fixed = DiffusionLayer::new(0, Activations::Linear, &network_fixed.sheaf)?;
network_fixed.sequential(vec![layer_fixed]);
let metrics_fixed = network_fixed.train_debug(
&training_data,
50,
false,
OptimKind::Adam,
0.01,
OptimizerParams::Else,
)?;
println!(
"Learned restrictions final loss: {}",
metrics_learned.final_loss
);
println!(
"Fixed restrictions final loss: {}",
metrics_fixed.final_loss
);
Ok(())
}
#[test]
fn test_multiple_diffusion_layers() -> Result<(), KohoError> {
let sheaf = create_triangle_sheaf()?;
let input = sheaf.get_k_cochain(0)?;
let target_data = vec![0.9f32, 0.8f32, 0.7f32];
let target = Matrix::from_slice(&target_data, 1, 3, Device::Cpu, DType::F32)?;
let training_data = vec![(input, target)];
let mut network = SheafNN::init(0, false, LossKind::MSE, sheaf);
let layer1 = DiffusionLayer::new(0, Activations::Softmax, &network.sheaf)?;
let layer2 = DiffusionLayer::new(0, Activations::Tanh, &network.sheaf)?;
let layer3 = DiffusionLayer::new(0, Activations::Sigmoid, &network.sheaf)?;
network.sequential(vec![layer1, layer2, layer3]);
let metrics = network.train_debug(
&training_data,
175,
false,
OptimKind::Adam,
0.15,
OptimizerParams::Else,
)?;
println!("Multi-layer network final loss: {}", metrics.final_loss);
let test_input = network.sheaf.get_k_cochain(0)?;
let output = network.forward(test_input, false)?;
assert_eq!(output.rows(), 1, "Output should have 3 vertices");
Ok(())
}
}