Files
foxhunt/ml/tests/kan_integration.rs
jgrusewski 57bae2cb68 fix(ml): OOM hardening + battle-test KAN/xLSTM/Diffusion models
Replace 8 unbounded Vec accumulation patterns with bounded VecDeque
across ensemble, PPO, DQN, Mamba2, and data pipeline code to prevent
OOM on RTX 3050 Ti (4GB VRAM) during live trading and extended training.

Key OOM fixes:
- Ensemble price/volatility history: Vec → VecDeque with O(1) eviction
- Data pipeline: MAX_FEATURES=500K cap (~512MB) prevents unbounded loading
- DQN replay buffer: full-array shuffle → HashSet random sampling (8MB → 256B)
- PPO loss histories: bounded VecDeque (cap 1K), eliminated batch.clone()
- Mamba2 scan: pre-allocated Vecs, explicit drop() after Tensor::cat
- Mamba2 training history: capped at 100, Tensor::randn replaces Vec→Tensor
- Mamba2 SSM reset: 2 unwrap() violations replaced with proper error handling

Battle-testing (19 new integration tests):
- KAN: 5 tests (forward, 50-epoch training 89.9% loss reduction, checkpoint)
- xLSTM: 7 tests (2D+3D forward, 30-epoch training 82% reduction, checkpoint)
- Diffusion: 7 tests (2D+3D forward, 20-epoch pipeline, checkpoint, validation)

Bonus: fix pre-existing cache test failure (match .dbn.zst files, graceful skip)

All 2390 lib tests pass, 0 new clippy errors.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-24 09:55:39 +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"));
}