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>
180 lines
5.4 KiB
Rust
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"));
|
|
}
|