use std::collections::{HashMap, HashSet}; use rand::{self, seq::IndexedRandom}; use serde::{Deserialize, Serialize}; use crate::{config::Config, trajectory::Action}; fn key(a: &str, b: &str) -> (String, String) { let a = a.to_string(); let b = b.to_string(); if a < b { (a, b) } else { (b, a) } } fn compute_sum( cand: &str, known_keys: &HashSet<&String>, map: &HashMap<(String, String), f32>, ) -> f32 { known_keys .iter() .filter_map(|k| map.get(&key(k, cand))) .sum() } 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 { similar: HashMap<(String, String), f32>, different: HashMap<(String, String), f32>, } #[derive(Debug)] pub enum Learning { SkipExtend, MoreExtend(Vec, String), SkipToMore(Vec, String), MoreToSkip(Vec, String), } impl Learner { pub fn prune(&mut self, valid: &HashSet) { self.similar .retain(|(a, b), _v| valid.contains(a) && valid.contains(b)); self.different .retain(|(a, b), _v| valid.contains(a) && valid.contains(b)); } pub fn learn(&mut self, config: &Config, learning: &Learning) { // global decay for (_k, v) in self.similar.iter_mut().chain(self.different.iter_mut()) { *v *= 1.0 - config.decay_rate; } // update match learning { Learning::SkipExtend => { // This is the least information carrying 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 damp // sub-linearly that update by a proxy of temporal distance. for (distance, e) in trajectory.iter().rev().enumerate() { let damp = ((distance + 1) as f32).powf(-0.5); let v = self.similar.entry(key(e, new)).or_default(); *v = (1.0 - config.learning_rate) * *v + config.learning_rate * damp; } } Learning::SkipToMore(trajectory, new) => { // We have learnt that `new` is different to everything in // `trajectory`, the strongest signal we have. for e in trajectory { let v = self.different.entry(key(e, new)).or_default(); *v = (1.0 - config.learning_rate) * *v + 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 e in trajectory { let v = self.different.entry(key(e, new)).or_default(); *v = (1.0 - config.learning_rate) * *v + config.learning_rate / damp; } } }; // global threshold drop self.similar = self .similar .extract_if(|_k, v| (*v).abs() > config.min_score_thresh) .collect(); self.different = self .different .extract_if(|_k, v| (*v).abs() > config.min_score_thresh) .collect(); } pub fn sample( &self, trajectory: &[String], action: &Action, candidates: &HashSet, temperature: f32, ) -> Option<(String, f32)> { let trajectory: HashSet<_> = trajectory.iter().collect(); let normalisation: f32 = f32::max(trajectory.len() as f32, 1.0); let candidates: Vec<_> = candidates .iter() .filter(|c| !trajectory.contains(c)) .collect(); if candidates.is_empty() { return None; } let items = candidates .iter() .map(|&c| { let sim = compute_sum(c, &trajectory, &self.similar); let dif = compute_sum(c, &trajectory, &self.different); let score = match action { Action::Skip => dif - sim, Action::More => sim - dif, }; (c, score / normalisation) }) .collect::>(); softmax_sample(&items, temperature) } }