use scirs2_core::ndarray::{Array, Array1, Array2, Dimension, IxDyn, ScalarOperand};
use scirs2_core::numeric::Float;
use scirs2_core::random::Random;
use std::fmt::Debug;
use crate::error::{OptimError, Result};
use crate::optimizers::Optimizer;
const EPSILON: f64 = 1e-12;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AddressingMode {
Content,
Location,
Hybrid,
}
#[derive(Debug, Clone)]
pub struct NtmConfig<A: Float + ScalarOperand + Debug> {
pub memory_slots: usize,
pub memory_width: usize,
pub learning_rate: A,
pub read_sharpness: A,
pub erase_gate: A,
pub addressing_mode: AddressingMode,
pub memory_weight: A,
pub gradient_weight: A,
pub seed: u64,
}
impl<A: Float + ScalarOperand + Debug> Default for NtmConfig<A> {
fn default() -> Self {
Self {
memory_slots: 32,
memory_width: 16,
learning_rate: A::from(0.01).unwrap_or_else(A::zero),
read_sharpness: A::from(1.0).unwrap_or_else(A::one),
erase_gate: A::from(0.5).unwrap_or_else(A::zero),
addressing_mode: AddressingMode::Hybrid,
memory_weight: A::from(0.3).unwrap_or_else(A::zero),
gradient_weight: A::from(0.7).unwrap_or_else(A::one),
seed: 42,
}
}
}
pub struct NtmOptimizer<A: Float + ScalarOperand + Debug> {
config: NtmConfig<A>,
memory: Array2<A>,
prev_read_weights: Array1<A>,
prev_write_weights: Array1<A>,
rng_seed: u64,
step_count: usize,
}
impl<A: Float + ScalarOperand + Debug> Debug for NtmOptimizer<A> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NtmOptimizer")
.field("memory_slots", &self.config.memory_slots)
.field("memory_width", &self.config.memory_width)
.field("learning_rate", &self.config.learning_rate)
.field("addressing_mode", &self.config.addressing_mode)
.field("rng_seed", &self.rng_seed)
.field("step_count", &self.step_count)
.finish()
}
}
impl<A: Float + ScalarOperand + Debug> NtmOptimizer<A> {
pub fn new(memory_slots: usize, memory_width: usize, learning_rate: A) -> Self {
let config = NtmConfig::<A> {
memory_slots,
memory_width,
learning_rate,
..NtmConfig::<A>::default()
};
Self::with_config(config)
}
pub fn with_config(config: NtmConfig<A>) -> Self {
let slots = config.memory_slots;
let width = config.memory_width;
let seed = config.seed;
let _rng: Random<scirs2_core::random::rngs::StdRng> = Random::seed(seed);
let memory = Array2::<A>::zeros((slots, width));
let prev_read = Array1::<A>::zeros(slots);
let prev_write = Array1::<A>::zeros(slots);
Self {
config,
memory,
prev_read_weights: prev_read,
prev_write_weights: prev_write,
rng_seed: seed,
step_count: 0,
}
}
pub fn with_read_sharpness(mut self, beta: A) -> Self {
self.config.read_sharpness = beta;
self
}
pub fn with_erase_gate(mut self, gate: A) -> Self {
self.config.erase_gate = gate;
self
}
pub fn with_addressing(mut self, mode: AddressingMode) -> Self {
self.config.addressing_mode = mode;
self
}
pub fn with_memory_weight(mut self, weight: A) -> Self {
self.config.memory_weight = weight;
self
}
pub fn with_gradient_weight(mut self, weight: A) -> Self {
self.config.gradient_weight = weight;
self
}
pub fn with_seed(mut self, seed: u64) -> Self {
self.config.seed = seed;
self.rng_seed = seed;
let _rng: Random<scirs2_core::random::rngs::StdRng> = Random::seed(seed);
self
}
pub fn config(&self) -> &NtmConfig<A> {
&self.config
}
pub fn memory(&self) -> &Array2<A> {
&self.memory
}
pub fn memory_mut(&mut self) -> &mut Array2<A> {
&mut self.memory
}
pub fn last_read_weights(&self) -> &Array1<A> {
&self.prev_read_weights
}
pub fn last_write_weights(&self) -> &Array1<A> {
&self.prev_write_weights
}
pub fn step_count(&self) -> usize {
self.step_count
}
pub fn reset(&mut self) {
self.memory.fill(A::zero());
self.prev_read_weights.fill(A::zero());
self.prev_write_weights.fill(A::zero());
self.step_count = 0;
let _rng: Random<scirs2_core::random::rngs::StdRng> = Random::seed(self.rng_seed);
}
fn validate_config(&self) -> Result<()> {
if self.config.memory_slots == 0 {
return Err(OptimError::InvalidConfig(
"NtmOptimizer: memory_slots must be > 0".to_string(),
));
}
if self.config.memory_width == 0 {
return Err(OptimError::InvalidConfig(
"NtmOptimizer: memory_width must be > 0".to_string(),
));
}
Ok(())
}
fn build_key<D: Dimension>(&self, gradients: &Array<A, D>) -> Array1<A> {
let w = self.config.memory_width;
let mut key = Array1::<A>::zeros(w);
if w == 0 {
return key;
}
let flat: Vec<A> = gradients.iter().copied().collect();
let n = flat.len();
if n == 0 {
return key;
}
if n >= w {
for j in 0..w {
let lo = (j * n) / w;
let hi = ((j + 1) * n) / w;
let hi_safe = hi.max(lo + 1).min(n);
let mut acc = A::zero();
let mut count: usize = 0;
for value in flat.iter().take(hi_safe).skip(lo) {
acc = acc + *value;
count += 1;
}
if count > 0 {
let denom = A::from(count).unwrap_or_else(A::one);
key[j] = acc / denom;
}
}
} else {
for (i, value) in flat.iter().enumerate() {
let j = (i * w) / n;
key[j] = key[j] + *value;
}
let mut counts = vec![0_usize; w];
for i in 0..n {
let j = (i * w) / n;
counts[j] += 1;
}
for (j, c) in counts.iter().enumerate() {
if *c > 1 {
let denom = A::from(*c).unwrap_or_else(A::one);
key[j] = key[j] / denom;
}
}
}
let w_a = A::from(w).unwrap_or_else(A::one);
let mean = key.iter().copied().fold(A::zero(), |acc, x| acc + x) / w_a;
let mut var = A::zero();
for value in key.iter() {
let d = *value - mean;
var = var + d * d;
}
var = var / w_a;
let eps = A::from(EPSILON).unwrap_or_else(A::epsilon);
let std = var.sqrt() + eps;
for v in key.iter_mut() {
*v = (*v - mean) / std;
}
key
}
fn cosine_similarity(a: &Array1<A>, b: &Array1<A>) -> A {
let eps = A::from(EPSILON).unwrap_or_else(A::epsilon);
let mut dot = A::zero();
let mut na = A::zero();
let mut nb = A::zero();
for (x, y) in a.iter().zip(b.iter()) {
dot = dot + (*x) * (*y);
na = na + (*x) * (*x);
nb = nb + (*y) * (*y);
}
let denom = na.sqrt() * nb.sqrt() + eps;
let sim = dot / denom;
let one = A::one();
let neg_one = -one;
if sim > one {
one
} else if sim < neg_one {
neg_one
} else {
sim
}
}
fn sharpened_softmax(sims: &Array1<A>, beta: A) -> Array1<A> {
let n = sims.len();
let mut out = Array1::<A>::zeros(n);
if n == 0 {
return out;
}
let mut max_val = sims[0] * beta;
for value in sims.iter().take(n).skip(1) {
let scaled = *value * beta;
if scaled > max_val {
max_val = scaled;
}
}
let mut sum = A::zero();
for (i, value) in sims.iter().enumerate() {
let exp_val = (*value * beta - max_val).exp();
out[i] = exp_val;
sum = sum + exp_val;
}
if sum > A::zero() {
for v in out.iter_mut() {
*v = *v / sum;
}
} else {
let denom = A::from(n).unwrap_or_else(A::one);
for v in out.iter_mut() {
*v = A::one() / denom;
}
}
out
}
fn shift_right(src: &Array1<A>) -> Array1<A> {
let n = src.len();
let mut out = Array1::<A>::zeros(n);
if n == 0 {
return out;
}
for i in 0..n {
let prev_index = (i + n - 1) % n;
out[i] = src[prev_index];
}
out
}
fn compute_attention(
&self,
content_weights: &Array1<A>,
prev_weights: &Array1<A>,
) -> Array1<A> {
match self.config.addressing_mode {
AddressingMode::Content => content_weights.clone(),
AddressingMode::Location => Self::shift_right(prev_weights),
AddressingMode::Hybrid => {
let n = content_weights.len();
let shifted = Self::shift_right(prev_weights);
let half = A::from(0.5).unwrap_or_else(|| A::one() / (A::one() + A::one()));
let mut mix = Array1::<A>::zeros(n);
for i in 0..n {
mix[i] = half * content_weights[i] + half * shifted[i];
}
let beta = self.config.read_sharpness;
let mut sharpened = Array1::<A>::zeros(n);
let mut sum = A::zero();
let zero = A::zero();
for i in 0..n {
let base = if mix[i] < zero { zero } else { mix[i] };
let powed = base.powf(beta);
sharpened[i] = powed;
sum = sum + powed;
}
if sum > A::zero() {
for v in sharpened.iter_mut() {
*v = *v / sum;
}
} else {
let denom = A::from(n).unwrap_or_else(A::one);
for v in sharpened.iter_mut() {
*v = A::one() / denom;
}
}
sharpened
}
}
}
fn read_from_memory(&self, weights: &Array1<A>) -> Array1<A> {
let w = self.config.memory_width;
let n = self.config.memory_slots;
let mut read = Array1::<A>::zeros(w);
for j in 0..w {
let mut acc = A::zero();
for i in 0..n {
acc = acc + weights[i] * self.memory[(i, j)];
}
read[j] = acc;
}
read
}
fn write_to_memory(&mut self, weights: &Array1<A>, key: &Array1<A>) {
let n = self.config.memory_slots;
let w = self.config.memory_width;
let erase = self.config.erase_gate;
let one = A::one();
for i in 0..n {
let w_i = weights[i];
for j in 0..w {
let e_j = erase * key[j];
let factor = one - w_i * e_j;
let current = self.memory[(i, j)];
self.memory[(i, j)] = current * factor + w_i * key[j];
}
}
}
fn tile_read_vector<D: Dimension>(
&self,
read: &Array1<A>,
params: &Array<A, D>,
) -> Result<Array<A, D>> {
let shape: Vec<usize> = params.shape().to_vec();
let total: usize = shape.iter().product();
let w = read.len();
let mut buf: Vec<A> = Vec::with_capacity(total);
if total == 0 {
} else if w == 0 {
for _ in 0..total {
buf.push(A::zero());
}
} else {
for i in 0..total {
buf.push(read[i % w]);
}
}
let dyn_arr = Array::<A, IxDyn>::from_shape_vec(IxDyn(&shape), buf).map_err(|err| {
OptimError::ComputationError(format!(
"NtmOptimizer: failed to reshape tiled read vector: {err}"
))
})?;
dyn_arr.into_dimensionality::<D>().map_err(|err| {
OptimError::DimensionMismatch(format!(
"NtmOptimizer: failed to project tiled read vector into target dimension: {err}"
))
})
}
}
impl<A, D> Optimizer<A, D> for NtmOptimizer<A>
where
A: Float + ScalarOperand + Debug,
D: Dimension,
{
fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
self.validate_config()?;
if params.shape() != gradients.shape() {
return Err(OptimError::DimensionMismatch(format!(
"NtmOptimizer::step: parameters have shape {:?} but gradients have shape {:?}",
params.shape(),
gradients.shape()
)));
}
let key = self.build_key(gradients);
let n = self.config.memory_slots;
let mut content_scores = Array1::<A>::zeros(n);
for i in 0..n {
let row = self.memory.row(i).to_owned();
content_scores[i] = Self::cosine_similarity(&key, &row);
}
let content_weights = Self::sharpened_softmax(&content_scores, self.config.read_sharpness);
let prev_read = self.prev_read_weights.clone();
let read_weights = self.compute_attention(&content_weights, &prev_read);
let read_vector = self.read_from_memory(&read_weights);
let write_weights = read_weights.clone();
self.write_to_memory(&write_weights, &key);
let tiled = self.tile_read_vector(&read_vector, params)?;
let g_w = self.config.gradient_weight;
let m_w = self.config.memory_weight;
let update = &(gradients * g_w) + &(&tiled * m_w);
let new_params = params - &(&update * self.config.learning_rate);
self.prev_read_weights = read_weights;
self.prev_write_weights = write_weights;
self.step_count = self.step_count.saturating_add(1);
Ok(new_params)
}
fn get_learning_rate(&self) -> A {
self.config.learning_rate
}
fn set_learning_rate(&mut self, learning_rate: A) {
self.config.learning_rate = learning_rate;
}
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::Array1;
#[test]
fn test_default_config_values() {
let cfg: NtmConfig<f64> = NtmConfig::default();
assert_eq!(cfg.memory_slots, 32);
assert_eq!(cfg.memory_width, 16);
assert!((cfg.learning_rate - 0.01).abs() < 1e-12);
assert!((cfg.read_sharpness - 1.0).abs() < 1e-12);
assert!((cfg.erase_gate - 0.5).abs() < 1e-12);
assert_eq!(cfg.addressing_mode, AddressingMode::Hybrid);
assert!((cfg.memory_weight - 0.3).abs() < 1e-12);
assert!((cfg.gradient_weight - 0.7).abs() < 1e-12);
assert_eq!(cfg.seed, 42);
}
#[test]
fn test_builder_pattern_chains() {
let opt: NtmOptimizer<f64> = NtmOptimizer::new(8, 4, 0.01)
.with_read_sharpness(2.5)
.with_erase_gate(0.25)
.with_addressing(AddressingMode::Content)
.with_memory_weight(0.4)
.with_gradient_weight(0.6)
.with_seed(7);
let cfg = opt.config();
assert!((cfg.read_sharpness - 2.5).abs() < 1e-12);
assert!((cfg.erase_gate - 0.25).abs() < 1e-12);
assert_eq!(cfg.addressing_mode, AddressingMode::Content);
assert!((cfg.memory_weight - 0.4).abs() < 1e-12);
assert!((cfg.gradient_weight - 0.6).abs() < 1e-12);
assert_eq!(cfg.seed, 7);
}
#[test]
fn test_new_initializes_memory_to_zero() {
let opt: NtmOptimizer<f64> = NtmOptimizer::new(4, 3, 0.01);
for &v in opt.memory().iter() {
assert_eq!(v, 0.0);
}
for &v in opt.last_read_weights().iter() {
assert_eq!(v, 0.0);
}
for &v in opt.last_write_weights().iter() {
assert_eq!(v, 0.0);
}
assert_eq!(opt.step_count(), 0);
}
#[test]
fn test_memory_dims_match_config() {
let opt: NtmOptimizer<f64> = NtmOptimizer::new(11, 5, 0.01);
assert_eq!(opt.memory().shape(), &[11, 5]);
assert_eq!(opt.last_read_weights().len(), 11);
assert_eq!(opt.last_write_weights().len(), 11);
}
#[test]
fn test_step_returns_same_shape_as_params() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(8, 4, 0.01);
let params = Array1::from_vec(vec![1.0, -1.0, 0.5, 2.0, -0.25]);
let grads = Array1::from_vec(vec![0.1, -0.2, 0.3, 0.0, -0.5]);
let next = opt.step(¶ms, &grads).expect("step failed");
assert_eq!(next.shape(), params.shape());
}
#[test]
fn test_step_changes_params() {
let mut opt: NtmOptimizer<f64> =
NtmOptimizer::new(8, 4, 0.1).with_addressing(AddressingMode::Content);
let params = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
let grads = Array1::from_vec(vec![0.5, -0.5, 0.25, -0.75]);
let next = opt.step(¶ms, &grads).expect("step failed");
let mut diff_total = 0.0_f64;
for (a, b) in next.iter().zip(params.iter()) {
diff_total += (a - b).abs();
}
assert!(
diff_total > 1e-6,
"non-zero gradient must update at least one parameter"
);
}
#[test]
fn test_zero_gradients_minimal_change() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(8, 4, 0.1);
let params = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
let grads = Array1::<f64>::zeros(4);
let next = opt.step(¶ms, &grads).expect("step failed");
for (a, b) in next.iter().zip(params.iter()) {
assert!(
(a - b).abs() < 1e-9,
"zero grads + zero memory must leave params unchanged (a={a}, b={b})"
);
}
}
#[test]
fn test_step_count_increments() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(8, 4, 0.01);
let params = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
let grads = Array1::from_vec(vec![0.1, -0.1, 0.2, -0.2]);
assert_eq!(opt.step_count(), 0);
let _ = opt.step(¶ms, &grads).expect("step 1 failed");
assert_eq!(opt.step_count(), 1);
let _ = opt.step(¶ms, &grads).expect("step 2 failed");
let _ = opt.step(¶ms, &grads).expect("step 3 failed");
assert_eq!(opt.step_count(), 3);
}
#[test]
fn test_addressing_mode_content_uses_pure_content() {
let mut opt: NtmOptimizer<f64> =
NtmOptimizer::new(5, 3, 0.01).with_addressing(AddressingMode::Content);
let params = Array1::from_vec(vec![0.0, 0.0, 0.0]);
let grads = Array1::from_vec(vec![1.0, -1.0, 0.5]);
let _ = opt.step(¶ms, &grads).expect("step failed");
let w = opt.last_read_weights();
for v in w.iter() {
assert!(
(*v - 0.2).abs() < 1e-6,
"Content mode with zero memory must yield uniform attention (got {v})"
);
}
}
#[test]
fn test_addressing_mode_location_shifts_weights() {
let mut opt: NtmOptimizer<f64> =
NtmOptimizer::new(5, 3, 0.01).with_addressing(AddressingMode::Location);
let _ = opt
.step(
&Array1::<f64>::zeros(3),
&Array1::from_vec(vec![1.0, 0.0, 0.0]),
)
.expect("warmup step failed");
let mut opt2: NtmOptimizer<f64> =
NtmOptimizer::new(5, 3, 0.01).with_addressing(AddressingMode::Content);
let _ = opt2
.step(
&Array1::<f64>::zeros(3),
&Array1::from_vec(vec![1.0, -1.0, 0.5]),
)
.expect("seed step failed");
let mut opt3 = opt2;
let _ = &mut opt3;
let mut opt_loc: NtmOptimizer<f64> =
NtmOptimizer::new(4, 2, 0.01).with_addressing(AddressingMode::Content);
{
let mem = opt_loc.memory_mut();
mem[(0, 0)] = 1.0;
mem[(0, 1)] = 0.0;
mem[(1, 0)] = 0.0;
mem[(1, 1)] = 1.0;
mem[(2, 0)] = -1.0;
mem[(2, 1)] = 0.0;
mem[(3, 0)] = 0.0;
mem[(3, 1)] = -1.0;
}
let _ = opt_loc
.step(
&Array1::<f64>::zeros(4),
&Array1::from_vec(vec![1.0, 0.0, -1.0, 0.0]),
)
.expect("content step failed");
let before = opt_loc.last_read_weights().clone();
let mut opt_loc2: NtmOptimizer<f64> =
NtmOptimizer::new(4, 2, 0.01).with_addressing(AddressingMode::Location);
{
let mem = opt_loc2.memory_mut();
for ((i, j), v) in opt_loc.memory().indexed_iter() {
mem[(i, j)] = *v;
}
}
let _ = opt_loc2
.step(
&Array1::<f64>::zeros(4),
&Array1::from_vec(vec![1.0, 0.0, -1.0, 0.0]),
)
.expect("warmup for location failed");
let shifted = NtmOptimizer::<f64>::shift_right(&before);
let mut a: Vec<f64> = before.iter().copied().collect();
let mut b: Vec<f64> = shifted.iter().copied().collect();
a.sort_by(|x, y| x.partial_cmp(y).expect("sort"));
b.sort_by(|x, y| x.partial_cmp(y).expect("sort"));
for (x, y) in a.iter().zip(b.iter()) {
assert!(
(x - y).abs() < 1e-12,
"Location-mode shift must permute the previous weights"
);
}
let max_diff = before
.iter()
.zip(shifted.iter())
.map(|(x, y)| (x - y).abs())
.fold(0.0_f64, f64::max);
assert!(
max_diff > 1e-9,
"Location shift on a non-uniform vector must yield a different vector"
);
}
#[test]
fn test_addressing_mode_hybrid_combines() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(6, 3, 0.01)
.with_addressing(AddressingMode::Hybrid)
.with_read_sharpness(2.0);
{
let mem = opt.memory_mut();
for i in 0..6 {
mem[(i, 0)] = i as f64 * 0.1;
mem[(i, 1)] = (5 - i) as f64 * 0.1;
mem[(i, 2)] = ((i as f64) - 2.5) * 0.1;
}
}
let params = Array1::<f64>::zeros(3);
let grads = Array1::from_vec(vec![0.5, -0.5, 0.5]);
let _ = opt.step(¶ms, &grads).expect("step failed");
let weights = opt.last_read_weights();
let mut total = 0.0_f64;
for v in weights.iter() {
assert!(*v >= -1e-12, "Hybrid weights must be non-negative");
total += *v;
}
assert!(
(total - 1.0).abs() < 1e-6,
"Hybrid weights must sum to one (got {total})"
);
}
#[test]
fn test_memory_updated_after_step() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(4, 3, 0.01);
let before = opt.memory().clone();
let params = Array1::from_vec(vec![1.0, -1.0, 0.5]);
let grads = Array1::from_vec(vec![0.5, 0.5, -0.5]);
let _ = opt.step(¶ms, &grads).expect("step failed");
let after = opt.memory();
let mut diff = 0.0_f64;
for (a, b) in after.iter().zip(before.iter()) {
diff += (a - b).abs();
}
assert!(
diff > 1e-9,
"memory must change after a step with non-zero key"
);
}
#[test]
fn test_read_weights_sum_to_one() {
let mut opt: NtmOptimizer<f64> =
NtmOptimizer::new(7, 4, 0.01).with_addressing(AddressingMode::Content);
let params = Array1::from_vec(vec![1.0, 2.0, -1.0, 0.5]);
let grads = Array1::from_vec(vec![0.3, -0.1, 0.4, -0.2]);
let _ = opt.step(¶ms, &grads).expect("step failed");
let total: f64 = opt.last_read_weights().iter().sum();
assert!(
(total - 1.0).abs() < 1e-6,
"read attention must sum to one (got {total})"
);
}
#[test]
fn test_reset_clears_memory_and_weights() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(4, 3, 0.01);
let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
let grads = Array1::from_vec(vec![0.5, -0.5, 0.25]);
let _ = opt.step(¶ms, &grads).expect("step 1 failed");
let _ = opt.step(¶ms, &grads).expect("step 2 failed");
assert_eq!(opt.step_count(), 2);
opt.reset();
assert_eq!(opt.step_count(), 0);
for v in opt.memory().iter() {
assert_eq!(*v, 0.0);
}
for v in opt.last_read_weights().iter() {
assert_eq!(*v, 0.0);
}
for v in opt.last_write_weights().iter() {
assert_eq!(*v, 0.0);
}
}
#[test]
fn test_seed_reproducibility() {
let mut a: NtmOptimizer<f64> = NtmOptimizer::new(6, 4, 0.05).with_seed(123);
let mut b: NtmOptimizer<f64> = NtmOptimizer::new(6, 4, 0.05).with_seed(123);
let params = Array1::from_vec(vec![1.0, -1.0, 0.5, 2.0, -0.3, 0.0]);
let grads = Array1::from_vec(vec![0.2, -0.4, 0.1, 0.0, -0.2, 0.3]);
for _ in 0..5 {
let na = a.step(¶ms, &grads).expect("a.step failed");
let nb = b.step(¶ms, &grads).expect("b.step failed");
for (x, y) in na.iter().zip(nb.iter()) {
assert!(
(x - y).abs() < 1e-12,
"seeded NTMs must produce identical outputs (a={x}, b={y})"
);
}
}
}
#[test]
fn test_get_set_learning_rate() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(4, 3, 0.05);
let lr_before =
<NtmOptimizer<f64> as Optimizer<f64, scirs2_core::ndarray::Ix1>>::get_learning_rate(
&opt,
);
assert!((lr_before - 0.05).abs() < 1e-12);
<NtmOptimizer<f64> as Optimizer<f64, scirs2_core::ndarray::Ix1>>::set_learning_rate(
&mut opt, 0.123,
);
let lr_after =
<NtmOptimizer<f64> as Optimizer<f64, scirs2_core::ndarray::Ix1>>::get_learning_rate(
&opt,
);
assert!((lr_after - 0.123).abs() < 1e-12);
}
#[test]
fn test_dimension_mismatch_errors() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(4, 3, 0.01);
let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
let grads = Array1::from_vec(vec![0.1, 0.2]); let err = opt.step(¶ms, &grads);
assert!(matches!(err, Err(OptimError::DimensionMismatch(_))));
}
#[test]
fn test_convergence_on_quadratic() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(8, 4, 0.05)
.with_addressing(AddressingMode::Content)
.with_gradient_weight(1.0)
.with_memory_weight(0.0)
.with_seed(2024);
let mut x = Array1::from_vec(vec![2.0]);
for _ in 0..100 {
let g = x.mapv(|v| 2.0 * v);
x = opt.step(&x, &g).expect("step failed");
}
assert!(
x[0].abs() < 2.0,
"convergence test must reduce |x| below initial value (got {})",
x[0]
);
assert!(
x[0].abs() < 1e-2,
"convergence must drive x close to zero (got {})",
x[0]
);
}
#[test]
fn test_zero_memory_slots_errors() {
let cfg = NtmConfig::<f64> {
memory_slots: 0,
..NtmConfig::<f64>::default()
};
let mut opt = NtmOptimizer::with_config(cfg);
let params = Array1::from_vec(vec![1.0, 2.0]);
let grads = Array1::from_vec(vec![0.1, 0.2]);
let err = opt.step(¶ms, &grads);
assert!(matches!(err, Err(OptimError::InvalidConfig(_))));
}
#[test]
fn test_zero_memory_width_errors() {
let cfg = NtmConfig::<f64> {
memory_width: 0,
..NtmConfig::<f64>::default()
};
let mut opt = NtmOptimizer::with_config(cfg);
let params = Array1::from_vec(vec![1.0, 2.0]);
let grads = Array1::from_vec(vec![0.1, 0.2]);
let err = opt.step(¶ms, &grads);
assert!(matches!(err, Err(OptimError::InvalidConfig(_))));
}
#[test]
fn test_shift_right_is_circular() {
let v = Array1::from_vec(vec![1.0_f64, 2.0, 3.0, 4.0]);
let s = NtmOptimizer::<f64>::shift_right(&v);
assert_eq!(s, Array1::from_vec(vec![4.0, 1.0, 2.0, 3.0]));
let empty: Array1<f64> = Array1::zeros(0);
let s_empty = NtmOptimizer::<f64>::shift_right(&empty);
assert_eq!(s_empty.len(), 0);
}
#[test]
fn test_cosine_similarity_basic() {
let a = Array1::from_vec(vec![1.0_f64, 0.0, 0.0]);
let b = Array1::from_vec(vec![1.0_f64, 0.0, 0.0]);
let s = NtmOptimizer::<f64>::cosine_similarity(&a, &b);
assert!((s - 1.0).abs() < 1e-6, "expected ~1.0, got {s}");
let c = Array1::from_vec(vec![-1.0_f64, 0.0, 0.0]);
let s2 = NtmOptimizer::<f64>::cosine_similarity(&a, &c);
assert!((s2 + 1.0).abs() < 1e-6, "expected ~-1.0, got {s2}");
let d = Array1::from_vec(vec![0.0_f64, 1.0, 0.0]);
let s3 = NtmOptimizer::<f64>::cosine_similarity(&a, &d);
assert!(s3.abs() < 1e-6, "expected ~0.0, got {s3}");
}
#[test]
fn test_step_2d_array_shapes_round_trip() {
let mut opt: NtmOptimizer<f64> = NtmOptimizer::new(8, 4, 0.01);
let params = scirs2_core::ndarray::Array2::<f64>::zeros((3, 5));
let grads = scirs2_core::ndarray::Array2::<f64>::from_elem((3, 5), 0.1);
let next = opt.step(¶ms, &grads).expect("2D step failed");
assert_eq!(next.shape(), &[3, 5]);
}
}