diff options
| author | tslil <tslil@posteo.de> | 2026-07-12 18:20:30 +0100 |
|---|---|---|
| committer | tslil <tslil@posteo.de> | 2026-07-13 19:56:12 +0100 |
| commit | e0241df8afde0483fc36cb2300b53127672486dc (patch) | |
| tree | 9fa2bec196498df5b5167b55e250bb96d56d96d5 /src/learner.rs | |
Graph basedgraph
Diffstat (limited to 'src/learner.rs')
| -rw-r--r-- | src/learner.rs | 161 |
1 files changed, 161 insertions, 0 deletions
diff --git a/src/learner.rs b/src/learner.rs new file mode 100644 index 0000000..91b24fe --- /dev/null +++ b/src/learner.rs @@ -0,0 +1,161 @@ +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>, String), + SkipToMore(Vec<String>, String), + MoreToSkip(Vec<String>, String), +} + +impl Learner { + pub fn prune(&mut self, valid: &HashSet<String>) { + 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<String>, + 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::<Vec<_>>(); + + softmax_sample(&items, temperature) + } +} |
