Files
foxhunt/crates/ml/tests/kan_integration.rs
jgrusewski 9c3d741a08 refactor: restructure repo — crates/, bin/, testing/ layout
Move 17 library crates into crates/, CLI binary into bin/fxt,
consolidate 10 test crates into testing/, split config crate
from deployment config files.

Root directory reduced from 38+ to ~17 directories.
All Cargo.toml paths and build.rs proto refs updated.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-25 11:56:00 +01:00

180 lines
5.4 KiB
Rust

//! KAN (Kolmogorov-Arnold Network) Integration Tests
//!
//! Validates the KAN trainable adapter end-to-end:
//! construction, forward pass, training loop, checkpoint save/load.
#![allow(unused_crate_dependencies)]
use candle_core::{Device, Tensor};
use ml::kan::config::KANConfig;
use ml::kan::trainable::KANTrainableAdapter;
use ml::training::unified_trainer::UnifiedTrainable;
fn small_kan_config() -> KANConfig {
KANConfig {
layer_widths: vec![10, 8, 4, 1],
grid_size: 3,
spline_order: 3,
learning_rate: 1e-3,
weight_decay: 1e-5,
grad_clip: 1.0,
}
}
#[test]
fn test_kan_construction() {
let config = small_kan_config();
let adapter = KANTrainableAdapter::new(config, &Device::Cpu);
assert!(
adapter.is_ok(),
"KAN construction failed: {:?}",
adapter.err()
);
let adapter = adapter.unwrap();
assert_eq!(adapter.model_type(), "KAN");
assert_eq!(adapter.get_step(), 0);
}
#[test]
fn test_kan_forward_pass() {
let config = small_kan_config();
let mut adapter = KANTrainableAdapter::new(config, &Device::Cpu).unwrap();
// [batch=4, input_dim=10]
let input = Tensor::randn(0f32, 1.0, &[4, 10], &Device::Cpu).unwrap();
let output = adapter.forward(&input);
assert!(output.is_ok(), "Forward failed: {:?}", output.err());
let output = output.unwrap();
assert_eq!(
output.dims(),
&[4, 1],
"Expected [4, 1], got {:?}",
output.dims()
);
}
#[test]
fn test_kan_training_loop_loss_decreases() {
let config = small_kan_config();
let mut adapter = KANTrainableAdapter::new(config, &Device::Cpu).unwrap();
// Synthetic regression: target = mean of inputs
let batch_size = 16;
let input_dim = 10;
let input = Tensor::randn(0f32, 0.5, &[batch_size, input_dim], &Device::Cpu).unwrap();
let target = input.mean_keepdim(1).unwrap();
let mut first_loss = None;
let mut last_loss = 0.0;
for epoch in 0..50 {
// Forward
let predictions = adapter.forward(&input).unwrap();
// Loss
let loss = adapter.compute_loss(&predictions, &target).unwrap();
let loss_val = loss.to_scalar::<f32>().unwrap() as f64;
if first_loss.is_none() {
first_loss = Some(loss_val);
}
last_loss = loss_val;
// Backward
let _grad_norm = adapter.backward(&loss).unwrap();
// Optimizer step
adapter.optimizer_step().unwrap();
adapter.zero_grad().unwrap();
if epoch % 10 == 0 {
println!("KAN epoch {}: loss = {:.6}", epoch, loss_val);
}
}
let first = first_loss.unwrap();
println!(
"KAN training: first_loss={:.6}, last_loss={:.6}, reduction={:.1}%",
first,
last_loss,
(1.0 - last_loss / first) * 100.0
);
assert!(
last_loss < first,
"Loss should decrease: first={}, last={}",
first,
last_loss
);
}
#[test]
fn test_kan_checkpoint_roundtrip() {
let config = small_kan_config();
let mut adapter = KANTrainableAdapter::new(config.clone(), &Device::Cpu).unwrap();
// Do a few training steps to change weights
let input = Tensor::randn(0f32, 1.0, &[4, 10], &Device::Cpu).unwrap();
let target = Tensor::randn(0f32, 0.1, &[4, 1], &Device::Cpu).unwrap();
for _ in 0..5 {
let pred = adapter.forward(&input).unwrap();
let loss = adapter.compute_loss(&pred, &target).unwrap();
adapter.backward(&loss).unwrap();
adapter.optimizer_step().unwrap();
}
// Save checkpoint
let tmp_dir = std::env::temp_dir().join("kan_test_checkpoint");
std::fs::create_dir_all(&tmp_dir).unwrap();
let checkpoint_path = tmp_dir.join("kan_ckpt");
let save_result = adapter.save_checkpoint(checkpoint_path.to_str().unwrap());
assert!(save_result.is_ok(), "Save failed: {:?}", save_result.err());
// Load into fresh adapter
let mut adapter2 = KANTrainableAdapter::new(config, &Device::Cpu).unwrap();
let load_result = adapter2.load_checkpoint(checkpoint_path.to_str().unwrap());
assert!(load_result.is_ok(), "Load failed: {:?}", load_result.err());
// Verify same predictions
let pred1 = adapter.forward(&input).unwrap();
let pred2 = adapter2.forward(&input).unwrap();
let diff = (pred1 - pred2)
.unwrap()
.abs()
.unwrap()
.sum_all()
.unwrap()
.to_scalar::<f32>()
.unwrap();
assert!(
diff < 1e-5,
"Checkpoint roundtrip predictions differ by {}",
diff
);
// Cleanup
let _ = std::fs::remove_dir_all(&tmp_dir);
}
#[test]
fn test_kan_metrics_collection() {
let config = small_kan_config();
let mut adapter = KANTrainableAdapter::new(config, &Device::Cpu).unwrap();
let input = Tensor::randn(0f32, 1.0, &[4, 10], &Device::Cpu).unwrap();
let target = Tensor::randn(0f32, 0.1, &[4, 1], &Device::Cpu).unwrap();
let pred = adapter.forward(&input).unwrap();
let loss = adapter.compute_loss(&pred, &target).unwrap();
adapter.backward(&loss).unwrap();
adapter.optimizer_step().unwrap();
let metrics = adapter.collect_metrics();
assert!(metrics.learning_rate > 0.0);
assert!(metrics.custom_metrics.contains_key("training_steps"));
assert!(metrics.custom_metrics.contains_key("grid_size"));
assert!(metrics.custom_metrics.contains_key("spline_order"));
}