Files
foxhunt/crates/ml/src/tft/variable_selection.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

284 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 = is_2d
.then(|| inputs.unsqueeze(1)) // [batch_size, 1, input_size]
.transpose()?;
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.0_f32, 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.0_f32; 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.0_f32; 8]; // 2 * 4
let inputs = Tensor::from_slice(&input_data, (2, 4), &device)?;
let context_data = vec![0.5_f32; 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(())
}
}