use std::collections::{HashMap, HashSet}; use rand::{self, seq::IndexedRandom}; use serde::{Deserialize, Serialize}; use crate::{config::Config, trajectory::Action}; fn compute_weight(past: &HashSet, now: &HashSet<&String>) -> f32 { let inter = now.iter().filter(|&&p| past.contains(p)).count(); if inter == 0 { return 0.0; } (inter as f32) / ((now.len() as f32).sqrt() * (past.len() as f32).sqrt()) } fn softmax_sample(items: &[(String, f32)], temperature: f32) -> Option<(String, f32)> { if items.is_empty() { return None; } let mut rng = rand::rng(); let max_val = items .iter() .map(|(_, v)| *v) .fold(f32::NEG_INFINITY, f32::max); let exp_values: Vec<_> = items .iter() .map(|(k, val)| (k, ((val - max_val) / temperature).exp())) .collect(); let sum: f32 = exp_values.iter().map(|(_, v)| *v).sum(); let norm_values: Vec<_> = exp_values.iter().map(|&(k, v)| (k, v / sum)).collect(); norm_values .choose_weighted(&mut rng, |item| item.1) .map(|p| (p.0.clone(), p.1)) .ok() } #[derive(Serialize, Deserialize)] enum Episode { Continue { group: HashSet, avoid: String, committed: String, }, Escape { from: HashSet, to: String, committed: String, }, } #[derive(Serialize, Deserialize, Default)] pub struct Intuition { history: Vec, } pub enum Learning { SkipToMore(Vec, String), MoreToSkip(Vec, String), } impl Intuition { pub fn learn(&mut self, learning: &Learning) { let stamp = chrono::Utc::now().format("%+").to_string(); match learning { Learning::SkipToMore(trajectory, new) => { self.history.push(Episode::Escape { from: trajectory.iter().map(String::clone).collect(), to: new.clone(), committed: stamp, }); } Learning::MoreToSkip(trajectory, new) => { self.history.push(Episode::Continue { group: trajectory.iter().map(String::clone).collect(), avoid: new.clone(), committed: stamp, }); } } } pub fn sample( &self, trajectory: &[String], action: &Action, candidates: &HashSet, config: &Config, ) -> Option<(String, f32)> { let decay = f32::powf( 0.5, 1.0 / (config.lookback_window_halflife_in_entries as f32), ); let trajectory: HashSet<_> = trajectory.iter().collect(); let candidates: Vec<_> = candidates .iter() .filter(|c| !trajectory.contains(c)) .collect(); if candidates.is_empty() { return None; } let mut items: HashMap = candidates.into_iter().map(|c| (c.clone(), 0.0)).collect(); for (dist, episode) in self.history.iter().rev().enumerate() { let kernel = decay.powf(dist as f32); match (episode, action) { (Episode::Escape { from, to, .. }, Action::Skip) => { let w = compute_weight(from, &trajectory); if let Some(to_w) = items.get_mut(to) { *to_w += config.max_weight * w * kernel; } for f in from { if let Some(from_weight) = items.get_mut(f) { *from_weight -= config.mid_weight * w * kernel; } } } (Episode::Escape { from, to, .. }, Action::More) => { let w = compute_weight(from, &trajectory); if trajectory.contains(to) { for f in from { if let Some(f_w) = items.get_mut(f) { *f_w -= config.low_weight * w * kernel; } } } } (Episode::Continue { group, avoid, .. }, Action::More) => { let w = compute_weight(group, &trajectory); for g in group { if let Some(g_w) = items.get_mut(g) { *g_w += config.max_weight * w * kernel; } } if let Some(a_w) = items.get_mut(avoid) { *a_w -= config.low_weight * w * kernel; } } (Episode::Continue { group, .. }, Action::Skip) => { let w = compute_weight(group, &trajectory); for g in group { if let Some(v) = items.get_mut(g) { *v -= config.high_weight * w * kernel; } } } } } let pairs: Vec<_> = items.into_iter().collect(); softmax_sample(&pairs, config.temperature) } }