use std::collections::{HashMap, HashSet}; use rand::{self, seq::IndexedRandom}; use serde::{Deserialize, Serialize}; use crate::{config::Config, trajectory::Action}; const DIM: usize = 16; // ----------------------------------------------------------------------------- // helpers fn dot(a: &[f32; DIM], b: &[f32; DIM]) -> f32 { a.iter().zip(b).map(|(x, y)| x * y).sum() } fn norm(a: &[f32; DIM]) -> f32 { dot(a, a).sqrt() } fn renorm(a: &mut [f32; DIM]) { let n = norm(a).max(1e-6); for x in a.iter_mut() { *x /= n; } } fn init_vec() -> [f32; DIM] { let mut v = [0.0; DIM]; for x in v.iter_mut() { *x = rand::random_range(-0.1..0.1); } renorm(&mut v); v } 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, Default)] pub struct Learner { embedding: HashMap, } impl Learner { pub fn prune(&mut self, valid: &HashSet) { self.embedding.retain(|k, _v| valid.contains(k)); } fn nudge(&mut self, target: &str, direction: &str, rate: f32) { let t = *self .embedding .entry(direction.into()) .or_insert_with(init_vec); let v = self.embedding.entry(target.into()).or_insert_with(init_vec); for i in 0..DIM { v[i] += rate * (t[i] - v[i]); } renorm(v); } } // ----------------------------------------------------------------------------- // core #[derive(Debug)] pub enum Learning { SkipExtend, MoreExtend(Vec, String), SkipToMore(Vec, String), MoreToSkip(Vec, String), } impl Learner { pub fn learn(&mut self, config: &Config, learning: &Learning) { match learning { Learning::SkipExtend => { /* This is the least informative case, it does not follow that the trajectory contains similar or different items, we may simply be seeking something in particular. */ } Learning::MoreExtend(trajectory, new) => { /* We're continuing a good run so `new` is similar to everything in `trajectory`, but to account for the possibilty that our mood has changed over the course of this streak we dampen this feedback. */ for (distance, similar) in trajectory.iter().rev().enumerate() { let damp = ((distance + 1) as f32).powf(-0.5); self.nudge(new, similar, config.learning_rate * damp); self.nudge(similar, new, config.learning_rate * damp * 0.5); // optional: symmetric, weaker } } Learning::SkipToMore(trajectory, new) => { /* We have learnt that `new` is different to everything in `trajectory`, the strongest signal we have. */ for different in trajectory { self.nudge(new, different, -config.learning_rate); } } Learning::MoreToSkip(trajectory, new) => { /* `new` could be different to everything in `trajectory`, or we simply changed our minds, so we have only weak evidence of difference. The positive coherence of trajectory was taken care of during MoreExtend above. */ let damp = (trajectory.len() + 1) as f32; for weakly_different in trajectory { self.nudge(new, weakly_different, -config.learning_rate / damp); } } } } fn mood(&self, trajectory: &[String]) -> Option<[f32; DIM]> { let mut m = [0.0; DIM]; let mut any = false; for (distance, historical) in trajectory.iter().rev().enumerate() { if let Some(v) = self.embedding.get(historical) { let damp = ((distance + 1) as f32).powf(-0.5); for i in 0..DIM { m[i] += damp * v[i]; } any = true; } } if !any || norm(&m) < 1e-6 { return None; } renorm(&mut m); Some(m) } pub fn sample( &self, trajectory: &[String], action: &Action, candidates: &HashSet, temperature: f32, ) -> Option<(String, f32)> { let seen: HashSet<_> = trajectory.iter().collect(); let candidates: Vec<_> = candidates.iter().filter(|c| !seen.contains(c)).collect(); if candidates.is_empty() { return None; } let Some(mood) = self.mood(trajectory) else { let items: Vec<_> = candidates.iter().map(|&c| (c, 0.0)).collect(); return softmax_sample(&items, temperature); }; let items: Vec<_> = candidates .iter() .map(|&c| { let score = match self.embedding.get(c.as_str()) { None => 0.0, Some(v) => { let cos = dot(v, &mood); match action { Action::More => cos, Action::Skip => -cos.abs(), } } }; (c, score) }) .collect(); softmax_sample(&items, temperature) } }