use crate::ineru::LongTermMemory;
use crate::types::Pattern;
use std::collections::VecDeque;
pub struct SurpriseGate {
threshold: f32,
recent_surprises: VecDeque<f32>,
calibration_window: usize,
mean_surprise: f32,
var_surprise: f32,
observation_count: usize,
}
impl SurpriseGate {
pub fn new(threshold: f32) -> Self {
Self {
threshold,
recent_surprises: VecDeque::with_capacity(100),
calibration_window: 100,
mean_surprise: 0.5,
var_surprise: 0.1,
observation_count: 0,
}
}
pub fn compute_surprise(&self, pattern: &Pattern, ltm: &LongTermMemory) -> f32 {
if ltm.is_empty() {
return 1.0;
}
let prediction = ltm.predict(&pattern.embedding);
let max_similarity = ltm.max_similarity(&pattern.embedding);
let raw_surprise = (1.0 - prediction) * 0.5 + (1.0 - max_similarity) * 0.5;
self.normalize_surprise(raw_surprise)
}
pub fn observe(&mut self, _pattern: &Pattern) {
self.observation_count += 1;
}
pub fn record_surprise(&mut self, surprise: f32) {
self.recent_surprises.push_back(surprise);
if self.recent_surprises.len() > self.calibration_window {
self.recent_surprises.pop_front();
}
self.observation_count += 1;
let n = self.observation_count as f32;
let delta = surprise - self.mean_surprise;
self.mean_surprise += delta / n;
let delta2 = surprise - self.mean_surprise;
self.var_surprise += delta * delta2;
}
fn normalize_surprise(&self, raw: f32) -> f32 {
if self.observation_count < 10 {
return raw;
}
let std = self.get_std().max(0.01);
let z_score = (raw - self.mean_surprise) / std;
1.0 / (1.0 + (-z_score).exp())
}
fn get_std(&self) -> f32 {
if self.observation_count > 1 {
(self.var_surprise / (self.observation_count as f32 - 1.0)).sqrt()
} else {
0.1
}
}
pub fn should_update(&self, surprise: f32) -> bool {
surprise > self.threshold
}
pub fn threshold(&self) -> f32 {
self.threshold
}
pub fn set_threshold(&mut self, threshold: f32) {
self.threshold = threshold.clamp(0.0, 1.0);
}
pub fn adaptive_threshold(&self) -> f32 {
let std = self.get_std();
(self.mean_surprise + std).min(1.0).max(0.1)
}
pub fn stats(&self) -> SurpriseStats {
SurpriseStats {
mean: self.mean_surprise,
std: self.get_std(),
threshold: self.threshold,
observation_count: self.observation_count,
}
}
}
#[derive(Debug, Clone)]
pub struct SurpriseStats {
pub mean: f32,
pub std: f32,
pub threshold: f32,
pub observation_count: usize,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{pattern_id, Embedding};
use std::collections::HashMap;
fn make_pattern(id: u8) -> Pattern {
let embedding = Embedding::new(vec![id as f32 / 255.0; 16]);
Pattern {
id: pattern_id(&[id]),
embedding,
metadata: HashMap::new(),
created_at: 1702656000000,
}
}
#[test]
fn test_surprise_gate_empty_memory() {
let gate = SurpriseGate::new(0.5);
let ltm = LongTermMemory::new(100, 16);
let pattern = make_pattern(1);
let surprise = gate.compute_surprise(&pattern, <m);
assert_eq!(surprise, 1.0); }
#[test]
fn test_surprise_gate_known_pattern() {
let gate = SurpriseGate::new(0.5);
let mut ltm = LongTermMemory::new(100, 16);
for i in 0..10 {
ltm.update(make_pattern(i)).unwrap();
}
let pattern = make_pattern(5);
let surprise = gate.compute_surprise(&pattern, <m);
assert!(surprise < 0.8);
}
#[test]
fn test_surprise_gate_novel_pattern() {
let gate = SurpriseGate::new(0.5);
let mut ltm = LongTermMemory::new(100, 16);
for i in 0..10 {
ltm.update(make_pattern(i)).unwrap();
}
let novel = Pattern {
id: pattern_id(&[255]),
embedding: Embedding::new(vec![1.0; 16]),
metadata: HashMap::new(),
created_at: 1702656000000,
};
let surprise = gate.compute_surprise(&novel, <m);
assert!(surprise > 0.3);
}
#[test]
fn test_adaptive_threshold() {
let mut gate = SurpriseGate::new(0.5);
for i in 0..50 {
let surprise = (i as f32) / 100.0;
gate.record_surprise(surprise);
}
let adaptive = gate.adaptive_threshold();
assert!(adaptive > 0.0 && adaptive < 1.0);
}
}