Files
foxhunt/ml/src/tft/variable_selection.rs
jgrusewski 987e5e6ac2 refactor(ml): remove 797 lines of commented-out code and disabled imports
Removed across 66 files:
- 49 instances of "// use crate::safe_operations; // DISABLED"
- 11 instances of "// use error_handling::{...}; // crate doesn't exist"
- 2 instances of "// use crate::Optimizer; // not available"
- 5 disabled test placeholder blocks (/* ... */) in ensemble/
- 1 disabled From impl in lib.rs (38 lines)
- 1 disabled test module in model.rs (113 lines)
- 1 disabled code block in integration/distillation.rs (41 lines)
- Various other disabled imports with explanation comments

All of this code references modules/crates that were removed during
prior refactoring waves and is preserved in git history. Removing it
reduces noise and makes the codebase easier to navigate.

1922 lib tests passing, compilation clean.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-20 19:43:47 +01:00

286 lines
10 KiB
Rust

//! Variable Selection Network for TFT
//!
//! Implements learnable feature selection using gated linear units and
//! soft feature selection weights for improved interpretability.
use std::collections::HashMap;
use candle_core::{Device, Module, Tensor};
use candle_nn::{linear, Linear, VarBuilder};
use super::GatedResidualNetwork;
use crate::MLError;
/// Variable Selection Network for feature importance learning
#[derive(Debug, Clone)]
pub struct VariableSelectionNetwork {
pub input_size: usize,
pub hidden_size: usize,
// Gated Linear Units for variable selection
flattened_grn: GatedResidualNetwork,
single_var_grns: Vec<GatedResidualNetwork>,
// Soft attention weights
attention_weights: Linear,
// Feature importance tracking
importance_scores: HashMap<usize, f64>,
device: Device,
}
impl VariableSelectionNetwork {
pub fn new(input_size: usize, hidden_size: usize, vs: VarBuilder<'_>) -> Result<Self, MLError> {
let device = vs.device().clone();
// Create GRN for flattened inputs
let flattened_grn =
GatedResidualNetwork::new(input_size, hidden_size, vs.pp("flattened_grn"))?;
// Create individual GRNs for each variable
let mut single_var_grns = Vec::new();
for i in 0..input_size {
let grn = GatedResidualNetwork::new(
1, // Single variable
hidden_size,
vs.pp(format!("single_var_grn_{}", i)),
)?;
single_var_grns.push(grn);
}
// Attention layer for variable selection
let attention_weights = linear(
hidden_size * input_size,
input_size,
vs.pp("attention_weights"),
)?;
Ok(Self {
input_size,
hidden_size,
flattened_grn,
single_var_grns,
attention_weights,
importance_scores: HashMap::new(),
device,
})
}
pub fn forward(
&mut self,
inputs: &Tensor,
context: Option<&Tensor>,
) -> Result<Tensor, MLError> {
let batch_size = inputs.dim(0)?;
let input_dims = inputs.dims();
// Normalize inputs to 3D format for uniform processing
// For 2D: create reshaped view once (reused across all variables)
// For 3D: use input reference directly to avoid 2.88MB clone per forward pass
let is_2d = input_dims.len() == 2;
let seq_len = if is_2d { 1 } else { input_dims[1] };
// Pre-reshape 2D inputs outside loop to avoid repeated operations
let reshaped_2d = if is_2d {
Some(inputs.unsqueeze(1)?) // [batch_size, 1, input_size]
} else {
None
};
if !is_2d && input_dims.len() != 3 {
return Err(MLError::InvalidInput(format!(
"Input must be 2D or 3D, got {:?}",
input_dims
)));
}
// Process individual variables
let mut var_outputs = Vec::new();
for (i, grn) in self.single_var_grns.iter_mut().enumerate() {
// Extract variable i from all time steps
// Use pre-reshaped tensor for 2D, or direct input reference for 3D
let var_data = if let Some(ref reshaped) = reshaped_2d {
reshaped.narrow(2, i, 1)? // [batch_size, 1, 1]
} else {
inputs.narrow(2, i, 1)? // [batch_size, seq_len, 1]
};
let var_flattened = var_data.flatten(1, 2)?; // [batch_size, seq_len]
let var_reshaped = var_flattened.unsqueeze(2)?; // [batch_size, seq_len, 1]
let var_flat_2d = var_reshaped.flatten(0, 1)?; // [batch_size * seq_len, 1]
let var_output = grn.forward(&var_flat_2d, context)?; // [batch_size * seq_len, hidden_size]
let var_output_3d = var_output.reshape((batch_size, seq_len, self.hidden_size))?;
var_outputs.push(var_output_3d);
}
// Stack variable outputs
let stacked_vars = Tensor::stack(&var_outputs, 3)?; // [batch_size, seq_len, hidden_size, input_size]
let vars_flattened = stacked_vars.flatten(2, 3)?; // [batch_size, seq_len, hidden_size * input_size]
// Compute attention weights for variable selection
let attention_input = vars_flattened.flatten(0, 1)?; // [batch_size * seq_len, hidden_size * input_size]
let raw_weights = self.attention_weights.forward(&attention_input)?; // [batch_size * seq_len, input_size]
let attention_weights = candle_nn::ops::softmax(&raw_weights, 1)?;
let attention_3d = attention_weights.reshape((batch_size, seq_len, self.input_size))?;
// Update importance scores
self.update_importance_scores(&attention_3d)?;
// Apply variable selection weights
let weighted_vars = self.apply_variable_selection(&stacked_vars, &attention_3d)?;
Ok(weighted_vars)
}
fn update_importance_scores(&mut self, attention_weights: &Tensor) -> Result<(), MLError> {
// Compute mean attention weights across batch and time
let mean_weights = attention_weights.mean_keepdim(0)?.mean_keepdim(1)?; // [1, 1, input_size]
let weights_vec = mean_weights.flatten_all()?.to_vec1::<f32>()?;
// Clear previous scores to prevent memory growth (HashMap maintains capacity but releases entries)
self.importance_scores.clear();
// Update importance scores
for (i, weight) in weights_vec.iter().copied().enumerate() {
self.importance_scores.insert(i, weight as f64);
}
Ok(())
}
fn apply_variable_selection(
&self,
stacked_vars: &Tensor,
attention_weights: &Tensor,
) -> Result<Tensor, MLError> {
// Expand attention weights to match stacked_vars dimensions
let expanded_weights = attention_weights.unsqueeze(2)?; // [batch_size, seq_len, 1, input_size]
let broadcast_weights = expanded_weights.broadcast_as(stacked_vars.shape())?;
// Apply weights
let weighted = (stacked_vars * &broadcast_weights)?;
// Sum over variables dimension
let selected = weighted.sum(3)?; // [batch_size, seq_len, hidden_size]
Ok(selected)
}
pub fn get_importance_scores(&self) -> Result<Vec<f64>, MLError> {
let mut scores = vec![0.0; self.input_size];
for (i, &score) in &self.importance_scores {
if *i < self.input_size {
scores[*i] = score;
}
}
// Normalize to sum to 1.0 if all scores are zero (uniform distribution)
let sum: f64 = scores.iter().sum();
if sum == 0.0 {
let uniform_score = 1.0 / self.input_size as f64;
scores.fill(uniform_score);
}
Ok(scores)
}
pub fn get_top_features(&self, k: usize) -> Vec<(usize, f64)> {
let mut features: Vec<(usize, f64)> = self
.importance_scores
.iter()
.map(|(&idx, &score)| (idx, score))
.collect();
features.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
features.truncate(k);
features
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::DType;
#[test]
fn test_variable_selection_network_creation() -> Result<(), MLError> {
let device = Device::Cpu;
let vs = VarBuilder::zeros(DType::F32, &device);
let vsn = VariableSelectionNetwork::new(10, 64, vs.pp("test"))?;
assert_eq!(vsn.input_size, 10);
assert_eq!(vsn.hidden_size, 64);
Ok(())
}
#[test]
fn test_variable_selection_forward_2d() -> Result<(), MLError> {
let device = Device::Cpu;
let vs = VarBuilder::zeros(DType::F32, &device);
let mut vsn = VariableSelectionNetwork::new(5, 32, vs.pp("test"))?;
// Create test input [batch_size=2, input_size=5]
let input_data = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let inputs = Tensor::from_slice(&input_data, (2, 5), &device)?;
let output = vsn.forward(&inputs, None)?;
// Output should have shape [batch_size=2, seq_len=1, hidden_size=32]
assert_eq!(output.dims(), &[2, 1, 32]);
Ok(())
}
#[test]
fn test_variable_selection_forward_3d() -> Result<(), MLError> {
let device = Device::Cpu;
let vs = VarBuilder::zeros(DType::F32, &device);
let mut vsn = VariableSelectionNetwork::new(3, 16, vs.pp("test"))?;
// Create test input [batch_size=2, seq_len=4, input_size=3]
let input_data = vec![1.0f32; 24]; // 2 * 4 * 3
let inputs = Tensor::from_slice(&input_data, (2, 4, 3), &device)?;
let output = vsn.forward(&inputs, None)?;
// Output should have shape [batch_size=2, seq_len=4, hidden_size=16]
assert_eq!(output.dims(), &[2, 4, 16]);
Ok(())
}
#[test]
fn test_variable_selection_with_context() -> Result<(), MLError> {
let device = Device::Cpu;
let vs = VarBuilder::zeros(DType::F32, &device);
let mut vsn = VariableSelectionNetwork::new(4, 24, vs.pp("test"))?;
// Create test input and context
let input_data = vec![1.0f32; 8]; // 2 * 4
let inputs = Tensor::from_slice(&input_data, (2, 4), &device)?;
let context_data = vec![0.5f32; 48]; // 2 * 24
let context = Tensor::from_slice(&context_data, (2, 24), &device)?;
let output = vsn.forward(&inputs, Some(&context))?;
// Output should have shape [batch_size=2, seq_len=1, hidden_size=24]
assert_eq!(output.dims(), &[2, 1, 24]);
Ok(())
}
#[test]
fn test_importance_scores() -> Result<(), MLError> {
let device = Device::Cpu;
let vs = VarBuilder::zeros(DType::F32, &device);
let vsn = VariableSelectionNetwork::new(5, 32, vs.pp("test"))?;
let scores = vsn.get_importance_scores()?;
assert_eq!(scores.len(), 5);
// Should sum to 1.0 (uniform distribution)
let sum: f64 = scores.iter().sum();
assert!((sum - 1.0).abs() < 1e-6);
Ok(())
}
}