use crate::ai::sdr::{SdrVector, SDR_WORDS, SDR_DIM, SDR_DENSITY}; const MIN_PERMANENCE: f64 = 0.2; const CONNECTED_PERMANENCE: f64 = 0.5; const PERMANENCE_INCREMENT: f64 = 0.05; const PERMANENCE_DECREMENT: f64 = 0.02; const MAX_SEGMENTS_PER_CELL: usize = 4; const PRUNE_INTERVAL: usize = 200; #[derive(Clone, Debug)] pub struct Synapse { pub bit_index: usize, pub permanence: f64, } impl Synapse { pub fn new(bit_index: usize, permanence: f64) -> Self { Synapse { bit_index, permanence } } pub fn connected(&self) -> bool { self.permanence >= CONNECTED_PERMANENCE } } #[derive(Clone, Debug)] pub struct DendriteSegment { pub synapses: Vec, } impl DendriteSegment { pub fn new() -> Self { DendriteSegment { synapses: Vec::new() } } pub fn overlap(&self, sdr: &SdrVector) -> u32 { self.synapses.iter() .filter(|s| s.connected()) .filter(|s| { let wi = s.bit_index / 64; let bi = s.bit_index % 64; (sdr.bits[wi] >> bi) & 1 == 1 }) .count() as u32 } pub fn reinforce(&mut self, sdr: &SdrVector) { for s in &mut self.synapses { let wi = s.bit_index / 64; let bi = s.bit_index % 64; if (sdr.bits[wi] >> bi) & 1 == 1 { s.permanence = (s.permanence + PERMANENCE_INCREMENT).min(1.0); } else { s.permanence = (s.permanence - PERMANENCE_DECREMENT).max(0.0); } } } pub fn reinforce_match_only(&mut self, sdr: &SdrVector) { for s in &mut self.synapses { let wi = s.bit_index / 64; let bi = s.bit_index % 64; if (sdr.bits[wi] >> bi) & 1 == 1 { s.permanence = (s.permanence + PERMANENCE_INCREMENT).min(1.0); } } } pub fn prune_weak(&mut self) { self.synapses.retain(|s| s.permanence >= MIN_PERMANENCE); } } pub struct TemporalCell { pub id: usize, pub segments: Vec, pub pattern: SdrVector, } impl TemporalCell { pub fn new(id: usize, pattern: SdrVector) -> Self { TemporalCell { id, segments: Vec::new(), pattern } } pub fn learn_segment(&mut self, input: &SdrVector) { if self.segments.len() >= MAX_SEGMENTS_PER_CELL { let mut scores: Vec<(usize, u32)> = self.segments.iter() .enumerate() .map(|(i, seg)| (i, seg.overlap(input))) .collect(); scores.sort_by(|a, b| b.1.cmp(&a.1)); scores.pop(); let idx = scores.last().map(|s| s.0).unwrap_or(0); self.segments.remove(idx); } let mut seg = DendriteSegment::new(); for bit in 0..SDR_DIM { let wi = bit / 64; let bi = bit % 64; if (input.bits[wi] >> bi) & 1 == 1 { seg.synapses.push(Synapse::new(bit, CONNECTED_PERMANENCE + 0.1)); } } if !seg.synapses.is_empty() { self.segments.push(seg); } } } pub struct TemporalMemory { pub cells: Vec, pub window: Vec, pub context_len: usize, pub step: usize, } impl TemporalMemory { pub fn new(capacity: usize, context_len: usize) -> Self { TemporalMemory { cells: Vec::with_capacity(capacity), window: Vec::new(), context_len, step: 0, } } pub fn get_or_create_cell(&mut self, pattern: &SdrVector) -> usize { if let Some((idx, _)) = self.cells.iter().enumerate() .find(|(_, c)| c.pattern.bits == pattern.bits) { return idx; } let id = self.cells.len(); self.cells.push(TemporalCell::new(id, pattern.clone())); id } pub fn reset(&mut self) { self.window.clear(); self.step = 0; } pub fn feed(&mut self, input: &SdrVector) -> (SdrVector, f64) { self.step += 1; let cell_id = self.get_or_create_cell(input); let prediction = if !self.window.is_empty() { let prev = &self.window[self.window.len() - 1]; self.predict_next(prev) } else { SdrVector::zero() }; let match_score = if prediction.popcount() > 0 { prediction.soft_overlap(input) } else { 0.0 }; if !self.window.is_empty() { let prev = self.window[self.window.len() - 1].clone(); for seg in &mut self.cells[cell_id].segments { if seg.overlap(&prev) > 0 { seg.reinforce(&prev); } } let mut has_prev_seg = false; for seg in &self.cells[cell_id].segments { if seg.overlap(&prev) >= 3 { has_prev_seg = true; break; } } if !has_prev_seg { self.cells[cell_id].learn_segment(&prev); } for c in &mut self.cells { if c.id == cell_id { continue; } for seg in &mut c.segments { if seg.overlap(&prev) >= 3 { seg.reinforce_match_only(&prev); } } } } self.window.push(input.clone()); while self.window.len() > self.context_len { self.window.remove(0); } if self.step % PRUNE_INTERVAL == 0 { self.prune(); } (prediction, match_score) } pub fn feed_no_learn(&mut self, input: &SdrVector) -> (SdrVector, f64) { self.step += 1; let _cell_id = self.get_or_create_cell(input); let prediction = if !self.window.is_empty() { let prev = &self.window[self.window.len() - 1]; self.predict_next(prev) } else { SdrVector::zero() }; let match_score = if prediction.popcount() > 0 { prediction.soft_overlap(input) } else { 0.0 }; self.window.push(input.clone()); while self.window.len() > self.context_len { self.window.remove(0); } if self.step % PRUNE_INTERVAL == 0 { self.prune(); } (prediction, match_score) } pub fn predict_next(&self, prev: &SdrVector) -> SdrVector { let mut pred = SdrVector::zero(); for c in &self.cells { let depolarized = c.segments.iter().any(|seg| seg.overlap(prev) >= 5); if depolarized { for i in 0..SDR_WORDS { pred.bits[i] |= c.pattern.bits[i]; } } } if pred.popcount() > 0 { let target = (SDR_DIM as f64 * SDR_DENSITY).ceil() as usize; let mut scored: Vec<(usize, u64)> = (0..SDR_DIM).filter(|&bit| { let wi = bit / 64; let bi = bit % 64; (pred.bits[wi] >> bi) & 1 == 1 }).map(|bit| { let score = (self.step as u64).wrapping_mul(bit as u64 + 1).reverse_bits(); (bit, score) }).collect(); scored.sort_by(|a, b| b.1.cmp(&a.1)); scored.truncate(target); let mut out = SdrVector::zero(); for &(bit, _) in &scored { out.bits[bit / 64] |= 1u64 << (bit % 64); } out } else { SdrVector::zero() } } pub fn learn_sequence(&mut self, prev: &SdrVector, next: &SdrVector) { let next_cell = self.get_or_create_cell(next); let mut has_prev_seg = false; for seg in &self.cells[next_cell].segments { if seg.overlap(prev) >= 3 { has_prev_seg = true; break; } } if !has_prev_seg { self.cells[next_cell].learn_segment(prev); } for seg in &mut self.cells[next_cell].segments { if seg.overlap(prev) > 0 { seg.reinforce(prev); } } } pub fn prune(&mut self) { for c in &mut self.cells { for seg in &mut c.segments { seg.prune_weak(); } c.segments.retain(|seg| { let connected = seg.synapses.iter().filter(|s| s.connected()).count(); connected >= 2 }); } } pub fn save(&self, path: &str) { use std::io::Write; if let Ok(mut f) = std::fs::File::create(path) { let n = self.cells.len() as u32; f.write_all(&n.to_le_bytes()).ok(); for c in &self.cells { f.write_all(&(c.id as u32).to_le_bytes()).ok(); for w in &c.pattern.bits { f.write_all(&w.to_le_bytes()).ok(); } f.write_all(&(c.segments.len() as u32).to_le_bytes()).ok(); for seg in &c.segments { f.write_all(&(seg.synapses.len() as u32).to_le_bytes()).ok(); for s in &seg.synapses { f.write_all(&(s.bit_index as u32).to_le_bytes()).ok(); f.write_all(&s.permanence.to_le_bytes()).ok(); } } } f.write_all(&(self.window.len() as u32).to_le_bytes()).ok(); for sdr in &self.window { for w in &sdr.bits { f.write_all(&w.to_le_bytes()).ok(); } } } } pub fn load(path: &str) -> Option { let data = std::fs::read(path).ok()?; let mut pos = 0usize; let n = u32::from_le_bytes(data[pos..pos+4].try_into().ok()?) as usize; pos += 4; let mut cells = Vec::with_capacity(n); for _ in 0..n { let id = u32::from_le_bytes(data[pos..pos+4].try_into().ok()?) as usize; pos += 4; let mut bits = [0u64; 128]; for w in bits.iter_mut() { *w = u64::from_le_bytes(data[pos..pos+8].try_into().ok()?); pos += 8; } let pattern = crate::ai::sdr::SdrVector { bits }; let seg_n = u32::from_le_bytes(data[pos..pos+4].try_into().ok()?) as usize; pos += 4; let mut segments = Vec::with_capacity(seg_n); for _ in 0..seg_n { let syn_n = u32::from_le_bytes(data[pos..pos+4].try_into().ok()?) as usize; pos += 4; let mut synapses = Vec::with_capacity(syn_n); for _ in 0..syn_n { let bi = u32::from_le_bytes(data[pos..pos+4].try_into().ok()?) as usize; pos += 4; let perm = f64::from_le_bytes(data[pos..pos+8].try_into().ok()?); pos += 8; synapses.push(Synapse::new(bi, perm)); } segments.push(DendriteSegment { synapses }); } cells.push(TemporalCell { id, segments, pattern }); } let wl = u32::from_le_bytes(data[pos..pos+4].try_into().ok()?) as usize; pos += 4; let mut window = Vec::with_capacity(wl); for _ in 0..wl { let mut bits = [0u64; 128]; for w in bits.iter_mut() { *w = u64::from_le_bytes(data[pos..pos+8].try_into().ok()?); pos += 8; } window.push(crate::ai::sdr::SdrVector { bits }); } Some(TemporalMemory { cells, window, context_len: 4, step: 0 }) } pub fn stats(&self) -> String { let total_segs: usize = self.cells.iter().map(|c| c.segments.len()).sum(); let total_syn: usize = self.cells.iter().flat_map(|c| c.segments.iter()).map(|s| s.synapses.len()).sum(); format!("cells={} segments={} synapses={} window={}", self.cells.len(), total_segs, total_syn, self.window.len()) } }