Files
rustytorch/crates/training/rtx-transformers/src/meta/mann.rs
T
2026-03-04 00:08:42 +00:00

754 lines
27 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Memory-Augmented Neural Networks (MANN) Implementation
//!
//! MANN extends neural networks with external memory for few-shot learning.
//! Key features:
//! - External memory matrix for storing past experiences
//! - Content-based addressing using cosine similarity
//! - LRU (Least Recently Used) memory management
//! - Integration with N-way K-shot learning tasks
//! - One-shot learning capability through memory
use super::*;
use rtx_tensor::{Tensor, Device, DType};
use rtx_autograd::{backward, clear_tape, NodeId};
use std::collections::{HashMap, VecDeque};
/// Configuration for Memory-Augmented Neural Networks
#[derive(Debug, Clone)]
pub struct MANNConfig {
/// Size of external memory matrix (number of memory slots)
pub memory_size: usize,
/// Dimension of each memory slot
pub memory_dim: usize,
/// Learning rate for controller network
pub learning_rate: f32,
/// Focus parameter for content-based addressing (higher = more focused)
pub focus_parameter: f32,
/// Input feature dimension
pub input_dim: usize,
/// Output dimension (number of classes)
pub output_dim: usize,
/// Controller network hidden dimension
pub controller_hidden_dim: usize,
}
impl Default for MANNConfig {
fn default() -> Self {
Self {
memory_size: 128,
memory_dim: 40,
learning_rate: 0.001,
focus_parameter: 1.0,
input_dim: 784,
output_dim: 5,
controller_hidden_dim: 200,
}
}
}
/// Statistics for MANN training and evaluation
#[derive(Debug, Clone, Default)]
pub struct MANNStats {
/// Number of episodes processed
pub episodes_processed: usize,
/// Average memory usage across episodes
pub avg_memory_usage: f32,
/// Number of memory reads performed
pub memory_reads: usize,
/// Number of memory writes performed
pub memory_writes: usize,
/// Average accuracy on query sets
pub avg_accuracy: f32,
/// Memory slots currently in use
pub memory_slots_used: usize,
}
/// Memory-Augmented Neural Network for Few-Shot Learning
#[derive(Debug)]
pub struct MANN {
/// External memory matrix (memory_size × memory_dim)
memory: Tensor,
/// Memory usage tracking (LRU order)
memory_usage: VecDeque<usize>,
/// Controller network weights
controller: ControllerNetwork,
/// Configuration
config: MANNConfig,
/// Training statistics
stats: MANNStats,
/// Device for computations
device: Device,
}
/// Controller network for memory operations
#[derive(Debug, Clone)]
pub struct ControllerNetwork {
/// Input to hidden layer weights
pub input_hidden: Tensor,
/// Hidden layer bias
pub hidden_bias: Tensor,
/// Hidden to output weights (for classification)
pub hidden_output: Tensor,
/// Output bias
pub output_bias: Tensor,
/// Key generation weights (for addressing)
pub hidden_key: Tensor,
/// Key bias
pub key_bias: Tensor,
}
/// Memory addressing weights and operations
#[derive(Debug, Clone)]
pub struct MemoryAddressing {
/// Read weights for each memory location
pub read_weights: Tensor,
/// Write weights for each memory location
pub write_weights: Tensor,
/// Key vector for content-based addressing
pub key: Tensor,
}
/// Result of memory read operation
#[derive(Debug, Clone)]
pub struct MemoryReadResult {
/// Read vector from memory
pub read_vector: Tensor,
/// Read weights used
pub read_weights: Tensor,
}
/// Result of memory write operation
#[derive(Debug, Clone)]
pub struct MemoryWriteResult {
/// Updated memory matrix
pub updated_memory: Tensor,
/// Write weights used
pub write_weights: Tensor,
}
impl MANN {
/// Create new MANN with given configuration
pub fn new(config: MANNConfig, device: &Device) -> Result<Self> {
// Initialize memory matrix with small random values
let memory_data: Vec<f32> = (0..config.memory_size * config.memory_dim)
.map(|i| (i as f32 * 0.01) % 0.2 - 0.1)
.collect();
let memory = Tensor::from_vec(
memory_data,
&[config.memory_size, config.memory_dim],
device.clone(),
)?;
// Initialize controller network
let controller = ControllerNetwork::new(&config, device)?;
// Initialize memory usage tracking (all slots available initially)
let memory_usage = VecDeque::from_iter(0..config.memory_size);
Ok(Self {
memory,
memory_usage,
controller,
config,
stats: MANNStats::default(),
device: device.clone(),
})
}
/// Process a single episode for few-shot learning
pub fn process_episode(&mut self, episode: &Episode, device: &Device) -> Result<MANNEpisodeResult> {
// Reset memory usage tracking for new episode
self.memory_usage = VecDeque::from_iter(0..self.config.memory_size);
let (support_x, support_y) = episode.support_batch(device)?;
let (query_x, query_y) = episode.query_batch(device)?;
// Process support set (write to memory)
let mut support_predictions = Vec::new();
let support_batch_size = support_x.shape()[0];
for i in 0..support_batch_size {
let input = support_x.slice(&[i..i+1, ..])?;
let label = support_y.slice(&[i..i+1])?;
// Forward pass through controller
let controller_output = self.controller.forward(&input)?;
// Generate key for content-based addressing
let key = self.controller.generate_key(&input)?;
// Write to memory (store input-label association)
let write_result = self.write_to_memory(&input, &key, device)?;
support_predictions.push(controller_output);
}
// Process query set (read from memory and classify)
let mut query_predictions = Vec::new();
let query_batch_size = query_x.shape()[0];
for i in 0..query_batch_size {
let input = query_x.slice(&[i..i+1, ..])?;
// Generate key for content-based addressing
let key = self.controller.generate_key(&input)?;
// Read from memory
let read_result = self.read_from_memory(&key, device)?;
// Classify using controller + memory read
let prediction = self.classify_with_memory(&input, &read_result.read_vector, device)?;
query_predictions.push(prediction);
}
// Compute accuracy
let accuracy = self.compute_accuracy(&query_predictions, &query_y, device)?;
// Update statistics
self.stats.episodes_processed += 1;
self.stats.avg_accuracy = (self.stats.avg_accuracy * (self.stats.episodes_processed - 1) as f32 + accuracy)
/ self.stats.episodes_processed as f32;
self.stats.memory_slots_used = self.config.memory_size - self.memory_usage.len();
Ok(MANNEpisodeResult {
support_predictions,
query_predictions,
accuracy,
memory_reads: query_batch_size,
memory_writes: support_batch_size,
})
}
/// Content-based addressing using cosine similarity
pub fn content_based_addressing(&self, key: &Tensor, device: &Device) -> Result<Tensor> {
let memory_size = self.config.memory_size;
let mut similarities = Vec::new();
// Compute cosine similarity between key and each memory slot
for i in 0..memory_size {
let memory_slot = self.memory.slice(&[i..i+1, ..])?;
let similarity = self.cosine_similarity(key, &memory_slot, device)?;
similarities.push(similarity);
}
// Convert to tensor
let similarities_tensor = Tensor::from_vec(similarities, &[memory_size], device.clone())?;
// Apply focus parameter and softmax
let focused = similarities_tensor.mul(&Tensor::scalar(self.config.focus_parameter, device.clone())?)?;
let weights = focused.softmax(0)?;
Ok(weights)
}
/// Compute cosine similarity between two vectors
pub fn cosine_similarity(&self, a: &Tensor, b: &Tensor, device: &Device) -> Result<f32> {
let dot_product = a.mul(b)?.sum(None)?.to_vec()?[0];
let norm_a = a.pow_tensor(&Tensor::scalar(2.0, device.clone())?)?.sum(None)?.sqrt()?.to_vec()?[0];
let norm_b = b.pow_tensor(&Tensor::scalar(2.0, device.clone())?)?.sum(None)?.sqrt()?.to_vec()?[0];
let similarity = dot_product / (norm_a * norm_b + 1e-8);
Ok(similarity)
}
/// Read from memory using content-based addressing
pub fn read_from_memory(&mut self, key: &Tensor, device: &Device) -> Result<MemoryReadResult> {
let read_weights = self.content_based_addressing(key, device)?;
// Weighted sum of memory slots
let read_vector = self.weighted_memory_read(&read_weights, device)?;
self.stats.memory_reads += 1;
Ok(MemoryReadResult {
read_vector,
read_weights,
})
}
/// Write to memory with LRU management
pub fn write_to_memory(&mut self, input: &Tensor, key: &Tensor, device: &Device) -> Result<MemoryWriteResult> {
let write_weights = self.content_based_addressing(key, device)?;
// Get least recently used memory slot
let lru_slot = self.get_lru_slot();
// Update memory at LRU slot
let input_data = input.to_vec()?;
let memory_data = self.memory.to_vec()?;
let mut new_memory_data = memory_data;
let start_idx = lru_slot * self.config.memory_dim;
let end_idx = start_idx + self.config.memory_dim.min(input_data.len());
for (i, &val) in input_data.iter().enumerate() {
if start_idx + i < new_memory_data.len() && i < self.config.memory_dim {
new_memory_data[start_idx + i] = val;
}
}
self.memory = Tensor::from_vec(
new_memory_data,
&[self.config.memory_size, self.config.memory_dim],
device.clone(),
)?;
// Update LRU tracking
self.update_lru_tracking(lru_slot);
self.stats.memory_writes += 1;
Ok(MemoryWriteResult {
updated_memory: self.memory.clone(),
write_weights,
})
}
/// Get least recently used memory slot
pub fn get_lru_slot(&self) -> usize {
self.memory_usage.front().copied().unwrap_or(0)
}
/// Update LRU tracking for given slot
pub fn update_lru_tracking(&mut self, slot: usize) {
// Remove slot from current position
self.memory_usage.retain(|&x| x != slot);
// Add to back (most recently used)
self.memory_usage.push_back(slot);
}
/// Perform weighted read from memory
pub fn weighted_memory_read(&self, weights: &Tensor, device: &Device) -> Result<Tensor> {
let memory_data = self.memory.to_vec()?;
let weights_data = weights.to_vec()?;
let memory_dim = self.config.memory_dim;
let memory_size = self.config.memory_size;
let mut read_data = vec![0.0f32; memory_dim];
for i in 0..memory_size {
let weight = weights_data[i];
let start_idx = i * memory_dim;
let end_idx = start_idx + memory_dim;
for (j, &mem_val) in memory_data[start_idx..end_idx].iter().enumerate() {
read_data[j] += weight * mem_val;
}
}
Ok(Tensor::from_vec(read_data, &[memory_dim], device.clone())?)
}
/// Classify using controller network and memory read
pub fn classify_with_memory(&self, input: &Tensor, memory_read: &Tensor, device: &Device) -> Result<Tensor> {
// Concatenate input with memory read
let input_data = input.to_vec()?;
let memory_data = memory_read.to_vec()?;
let mut combined_data = input_data;
combined_data.extend(memory_data);
let combined_input = Tensor::from_vec(
combined_data,
&[1, input.shape()[1] + memory_read.shape()[0]],
device.clone(),
)?;
// Forward pass through controller for classification
self.controller.forward(&combined_input)
}
/// Compute accuracy of predictions
pub fn compute_accuracy(&self, predictions: &[Tensor], targets: &Tensor, device: &Device) -> Result<f32> {
let targets_data = targets.to_vec()?;
let mut correct = 0;
let mut total = 0;
for (pred, &target) in predictions.iter().zip(targets_data.iter()) {
let pred_data = pred.to_vec()?;
let predicted_class = pred_data
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.total_cmp(b))
.map(|(idx, _)| idx)
.unwrap_or(0);
if predicted_class == target as usize {
correct += 1;
}
total += 1;
}
Ok(correct as f32 / total as f32)
}
/// Get current statistics
pub fn get_stats(&self) -> MANNStats {
self.stats.clone()
}
/// Clear memory and reset usage tracking
pub fn clear_memory(&mut self) -> Result<()> {
let memory_data = vec![0.0f32; self.config.memory_size * self.config.memory_dim];
self.memory = Tensor::from_vec(
memory_data,
&[self.config.memory_size, self.config.memory_dim],
self.device.clone(),
)?;
self.memory_usage = VecDeque::from_iter(0..self.config.memory_size);
Ok(())
}
}
/// Result of processing an episode with MANN
#[derive(Debug, Clone)]
pub struct MANNEpisodeResult {
/// Predictions on support set
pub support_predictions: Vec<Tensor>,
/// Predictions on query set
pub query_predictions: Vec<Tensor>,
/// Accuracy on query set
pub accuracy: f32,
/// Number of memory reads performed
pub memory_reads: usize,
/// Number of memory writes performed
pub memory_writes: usize,
}
impl ControllerNetwork {
/// Create new controller network
pub fn new(config: &MANNConfig, device: &Device) -> Result<Self> {
let input_dim = config.input_dim + config.memory_dim; // Input + memory read
let hidden_dim = config.controller_hidden_dim;
// Initialize weights with Xavier initialization
let input_hidden = Self::xavier_init(input_dim, hidden_dim, device)?;
let hidden_bias = Tensor::zeros(&[hidden_dim], device.clone())?;
let hidden_output = Self::xavier_init(hidden_dim, config.output_dim, device)?;
let output_bias = Tensor::zeros(&[config.output_dim], device.clone())?;
let hidden_key = Self::xavier_init(hidden_dim, config.memory_dim, device)?;
let key_bias = Tensor::zeros(&[config.memory_dim], device.clone())?;
Ok(Self {
input_hidden,
hidden_bias,
hidden_output,
output_bias,
hidden_key,
key_bias,
})
}
/// Xavier weight initialization
fn xavier_init(input_dim: usize, output_dim: usize, device: &Device) -> Result<Tensor> {
let limit = (6.0f32 / (input_dim + output_dim) as f32).sqrt();
let data: Vec<f32> = (0..input_dim * output_dim)
.map(|i| ((i as f32 * 0.1234) % 2.0 - 1.0) * limit)
.collect();
Ok(Tensor::from_vec(data, &[input_dim, output_dim], device.clone())?)
}
/// Forward pass through controller network
pub fn forward(&self, input: &Tensor) -> Result<Tensor> {
let hidden = input.matmul(&self.input_hidden)?.add(&self.hidden_bias)?.relu()?;
let output = hidden.matmul(&self.hidden_output)?.add(&self.output_bias)?;
Ok(output)
}
/// Generate key vector for memory addressing
pub fn generate_key(&self, input: &Tensor) -> Result<Tensor> {
let hidden = input.matmul(&self.input_hidden)?.add(&self.hidden_bias)?.relu()?;
let key = hidden.matmul(&self.hidden_key)?.add(&self.key_bias)?.tanh()?;
Ok(key)
}
}
#[cfg(all(test, feature = "disabled_tests"))]
mod tests {
use super::*;
#[test]
fn test_mann_config_creation() {
let config = MANNConfig::default();
assert_eq!(config.memory_size, 128);
assert_eq!(config.memory_dim, 40);
assert_eq!(config.focus_parameter, 1.0);
}
#[test]
fn test_mann_creation() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
memory_size: 10,
memory_dim: 5,
input_dim: 8,
output_dim: 3,
..Default::default()
};
let mann = MANN::new(config, &device).unwrap();
assert_eq!(mann.memory.shape(), &[10, 5]);
assert_eq!(mann.memory_usage.len(), 10);
}
#[test]
fn test_cosine_similarity_computation() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig::default();
let mann = MANN::new(config, &device).unwrap();
let a = Tensor::from_vec(vec![1.0, 0.0], &[1, 2], device.clone()).unwrap();
let b = Tensor::from_vec(vec![0.0, 1.0], &[1, 2], device.clone()).unwrap();
let c = Tensor::from_vec(vec![1.0, 0.0], &[1, 2], device.clone()).unwrap();
let sim_ab = mann.cosine_similarity(&a, &b, &device).unwrap();
let sim_ac = mann.cosine_similarity(&a, &c, &device).unwrap();
assert!((sim_ab - 0.0).abs() < 1e-6); // Orthogonal vectors
assert!((sim_ac - 1.0).abs() < 1e-6); // Identical vectors
}
#[test]
fn test_content_based_addressing() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
memory_size: 3,
memory_dim: 4,
focus_parameter: 1.0,
..Default::default()
};
let mann = MANN::new(config, &device).unwrap();
let key = Tensor::from_vec(vec![1.0, 0.0, 0.0, 0.0], &[1, 4], device.clone()).unwrap();
let weights = mann.content_based_addressing(&key, &device).unwrap();
assert_eq!(weights.shape(), &[3]);
let weights_data = weights.to_vec().unwrap();
let sum: f32 = weights_data.iter().sum();
assert!((sum - 1.0).abs() < 1e-6); // Should sum to 1 (softmax)
}
#[test]
fn test_memory_read_operation() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
memory_size: 5,
memory_dim: 3,
..Default::default()
};
let mut mann = MANN::new(config, &device).unwrap();
let key = Tensor::from_vec(vec![1.0, 2.0, 3.0], &[1, 3], device.clone()).unwrap();
let result = mann.read_from_memory(&key, &device).unwrap();
assert_eq!(result.read_vector.shape(), &[3]);
assert_eq!(result.read_weights.shape(), &[5]);
assert_eq!(mann.stats.memory_reads, 1);
}
#[test]
fn test_memory_write_operation() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
memory_size: 4,
memory_dim: 6,
..Default::default()
};
let mut mann = MANN::new(config, &device).unwrap();
let input = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[1, 6], device.clone()).unwrap();
let key = Tensor::from_vec(vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0], &[1, 6], device.clone()).unwrap();
let old_memory_usage_len = mann.memory_usage.len();
let result = mann.write_to_memory(&input, &key, &device).unwrap();
assert_eq!(result.updated_memory.shape(), &[4, 6]);
assert_eq!(result.write_weights.shape(), &[4]);
assert_eq!(mann.stats.memory_writes, 1);
assert_eq!(mann.memory_usage.len(), old_memory_usage_len); // LRU tracking updated
}
#[test]
fn test_lru_memory_management() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
memory_size: 3,
memory_dim: 2,
..Default::default()
};
let mut mann = MANN::new(config, &device).unwrap();
// Initially, slot 0 should be LRU
assert_eq!(mann.get_lru_slot(), 0);
// Use slot 0, should move to back
mann.update_lru_tracking(0);
assert_eq!(mann.get_lru_slot(), 1);
// Use slot 2, should move to back
mann.update_lru_tracking(2);
assert_eq!(mann.get_lru_slot(), 1);
}
#[test]
fn test_weighted_memory_read() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
memory_size: 2,
memory_dim: 3,
..Default::default()
};
let mann = MANN::new(config, &device).unwrap();
let weights = Tensor::from_vec(vec![0.7, 0.3], &[2], device.clone()).unwrap();
let result = mann.weighted_memory_read(&weights, &device).unwrap();
assert_eq!(result.shape(), &[3]);
}
#[test]
fn test_controller_network_creation() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
input_dim: 10,
memory_dim: 8,
controller_hidden_dim: 16,
output_dim: 5,
..Default::default()
};
let controller = ControllerNetwork::new(&config, &device).unwrap();
assert_eq!(controller.input_hidden.shape(), &[18, 16]); // input_dim + memory_dim, hidden_dim
assert_eq!(controller.hidden_bias.shape(), &[16]);
assert_eq!(controller.hidden_output.shape(), &[16, 5]);
assert_eq!(controller.output_bias.shape(), &[5]);
}
#[test]
fn test_controller_forward_pass() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
input_dim: 4,
memory_dim: 3,
controller_hidden_dim: 8,
output_dim: 2,
..Default::default()
};
let controller = ControllerNetwork::new(&config, &device).unwrap();
let input = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0], &[1, 7], device.clone()).unwrap();
let output = controller.forward(&input).unwrap();
assert_eq!(output.shape(), &[1, 2]);
}
#[test]
fn test_controller_key_generation() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
input_dim: 5,
memory_dim: 4,
controller_hidden_dim: 10,
output_dim: 3,
..Default::default()
};
let controller = ControllerNetwork::new(&config, &device).unwrap();
let input = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0], &[1, 9], device.clone()).unwrap();
let key = controller.generate_key(&input).unwrap();
assert_eq!(key.shape(), &[1, 4]);
// Check that tanh activation was applied (values should be in [-1, 1])
let key_data = key.to_vec().unwrap();
for &val in &key_data {
assert!(val >= -1.0 && val <= 1.0);
}
}
#[test]
fn test_mann_episode_processing() {
let device = Device::cuda(0).unwrap_or(Device::default());
// Create synthetic dataset and episode
let dataset = FewShotDataset::synthetic(3, 10, 20, &device).unwrap();
let episode = dataset.sample_episode(2, 2, 3, &device).unwrap();
let config = MANNConfig {
memory_size: 16,
memory_dim: 8,
input_dim: 20,
output_dim: 2,
controller_hidden_dim: 32,
..Default::default()
};
let mut mann = MANN::new(config, &device).unwrap();
let result = mann.process_episode(&episode, &device).unwrap();
assert_eq!(result.support_predictions.len(), 4); // 2 classes × 2 shots
assert_eq!(result.query_predictions.len(), 6); // 2 classes × 3 queries
assert!(result.accuracy >= 0.0 && result.accuracy <= 1.0);
assert_eq!(result.memory_reads, 6);
assert_eq!(result.memory_writes, 4);
assert_eq!(mann.stats.episodes_processed, 1);
}
#[test]
fn test_mann_accuracy_computation() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig::default();
let mann = MANN::new(config, &device).unwrap();
// Perfect predictions
let pred1 = Tensor::from_vec(vec![0.9, 0.1], &[2], device.clone()).unwrap();
let pred2 = Tensor::from_vec(vec![0.2, 0.8], &[2], device.clone()).unwrap();
let predictions = vec![pred1, pred2];
let targets = Tensor::from_vec(vec![0.0, 1.0], &[2], device.clone()).unwrap();
let accuracy = mann.compute_accuracy(&predictions, &targets, &device).unwrap();
assert!((accuracy - 1.0).abs() < 1e-6);
// Imperfect predictions
let pred1 = Tensor::from_vec(vec![0.1, 0.9], &[2], device.clone()).unwrap(); // Wrong
let pred2 = Tensor::from_vec(vec![0.2, 0.8], &[2], device.clone()).unwrap(); // Correct
let predictions = vec![pred1, pred2];
let accuracy = mann.compute_accuracy(&predictions, &targets, &device).unwrap();
assert!((accuracy - 0.5).abs() < 1e-6);
}
#[test]
fn test_mann_memory_clearing() {
let device = Device::cuda(0).unwrap_or(Device::default());
let config = MANNConfig {
memory_size: 5,
memory_dim: 4,
..Default::default()
};
let mut mann = MANN::new(config, &device).unwrap();
// Write something to memory first
let input = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], &[1, 4], device.clone()).unwrap();
let key = Tensor::from_vec(vec![1.0, 1.0, 1.0, 1.0], &[1, 4], device.clone()).unwrap();
let _ = mann.write_to_memory(&input, &key, &device).unwrap();
// Memory should have non-zero values
let memory_data = mann.memory.to_vec().unwrap();
let has_nonzero = memory_data.iter().any(|&x| x.abs() > 1e-6);
assert!(has_nonzero);
// Clear memory
mann.clear_memory().unwrap();
// Memory should be all zeros
let cleared_memory_data = mann.memory.to_vec().unwrap();
let all_zero = cleared_memory_data.iter().all(|&x| x.abs() < 1e-6);
assert!(all_zero);
// Memory usage should be reset
assert_eq!(mann.memory_usage.len(), 5);
assert_eq!(mann.get_lru_slot(), 0);
}
}