series.rs
8.4 kB · rust · 308 lines
1use crate::word::{observe, series_value, CODES};2use std::collections::HashMap;34#[derive(Clone, Copy, PartialEq, Eq, Debug)]5pub struct Frac {6 num: i128,7 den: i128,8}910fn gcd(a: i128, b: i128) -> i128 {11 let (mut a, mut b) = (a.abs(), b.abs());12 while b != 0 {13 let t = a % b;14 a = b;15 b = t;16 }17 if a == 0 {18 119 } else {20 a21 }22}2324impl Frac {25 pub fn new(num: i128, den: i128) -> Frac {26 assert!(den != 0, "a fraction needs a nonzero denominator");27 let sign = if den < 0 { -1 } else { 1 };28 let g = gcd(num, den);29 Frac {30 num: sign * num / g,31 den: sign * den / g,32 }33 }3435 pub fn int(value: i64) -> Frac {36 Frac {37 num: value as i128,38 den: 1,39 }40 }4142 pub fn zero() -> Frac {43 Frac { num: 0, den: 1 }44 }4546 pub fn is_zero(&self) -> bool {47 self.num == 048 }4950 pub fn add(&self, other: &Frac) -> Frac {51 let left = self.num.checked_mul(other.den).expect("no overflow");52 let right = other.num.checked_mul(self.den).expect("no overflow");53 Frac::new(54 left.checked_add(right).expect("no overflow"),55 self.den.checked_mul(other.den).expect("no overflow"),56 )57 }5859 pub fn mul(&self, other: &Frac) -> Frac {60 Frac::new(61 self.num.checked_mul(other.num).expect("no overflow"),62 self.den.checked_mul(other.den).expect("no overflow"),63 )64 }6566 pub fn neg(&self) -> Frac {67 Frac {68 num: -self.num,69 den: self.den,70 }71 }7273 pub fn div(&self, other: &Frac) -> Frac {74 assert!(!other.is_zero(), "no division by zero");75 Frac::new(76 self.num.checked_mul(other.den).expect("no overflow"),77 self.den.checked_mul(other.num).expect("no overflow"),78 )79 }8081 pub fn integer(&self) -> i64 {82 assert!(self.den == 1, "the entry is not an integer");83 self.num as i6484 }8586 pub fn parts(&self) -> (i128, i128) {87 (self.num, self.den)88 }8990 pub fn below(&self, other: &Frac) -> bool {91 self.num * other.den < other.num * self.den92 }93}9495pub struct Table {96 values: HashMap<Vec<u8>, [i64; 4]>,97}9899impl Table {100 pub fn new() -> Table {101 Table {102 values: HashMap::new(),103 }104 }105106 pub fn get(&mut self, word: &[u8], which: usize) -> i64 {107 if let Some(found) = self.values.get(word) {108 return found[which];109 }110 let obs = observe(word);111 let row = [112 series_value(&obs, 0),113 series_value(&obs, 1),114 series_value(&obs, 2),115 series_value(&obs, 3),116 ];117 self.values.insert(word.to_vec(), row);118 row[which]119 }120}121122pub struct Rep {123 pub basis: Vec<Vec<u8>>,124 pub matrices: HashMap<u8, Vec<Vec<i64>>>,125 pub lambda: Vec<i64>,126 pub gamma: Vec<i64>,127}128129fn suffixes() -> Vec<Vec<u8>> {130 let mut out = vec![Vec::new()];131 for &a in CODES.iter() {132 out.push(vec![a]);133 }134 for &a in CODES.iter() {135 for &b in CODES.iter() {136 out.push(vec![a, b]);137 }138 }139 out140}141142struct Space {143 rows: Vec<Vec<Frac>>,144 echelon: Vec<(usize, Vec<Frac>, Vec<Frac>)>,145}146147impl Space {148 fn new() -> Space {149 Space {150 rows: Vec::new(),151 echelon: Vec::new(),152 }153 }154155 fn reduce(&self, target: &[Frac]) -> (Vec<Frac>, Vec<Frac>) {156 let mut rest: Vec<Frac> = target.to_vec();157 let mut coefficients = vec![Frac::zero(); self.rows.len()];158 for (pivot, row, combination) in self.echelon.iter() {159 if rest[*pivot].is_zero() {160 continue;161 }162 let factor = rest[*pivot].div(&row[*pivot]);163 for (slot, value) in rest.iter_mut().zip(row.iter()) {164 *slot = slot.add(&factor.mul(value).neg());165 }166 for (slot, value) in coefficients.iter_mut().zip(combination.iter()) {167 *slot = slot.add(&factor.mul(value));168 }169 }170 (coefficients, rest)171 }172173 fn insert(&mut self, target: Vec<Frac>) -> bool {174 let (coefficients, rest) = self.reduce(&target);175 let pivot = rest.iter().position(|value| !value.is_zero());176 match pivot {177 None => false,178 Some(at) => {179 let mut combination: Vec<Frac> = coefficients.iter().map(|c| c.neg()).collect();180 combination.push(Frac::int(1));181 for (_, _, old) in self.echelon.iter_mut() {182 old.push(Frac::zero());183 }184 self.rows.push(target);185 self.echelon.push((at, rest, combination));186 true187 }188 }189 }190}191192fn row_of(word: &[u8], which: usize, suffix: &[Vec<u8>], table: &mut Table) -> Vec<Frac> {193 suffix194 .iter()195 .map(|tail| {196 let mut full = word.to_vec();197 full.extend_from_slice(tail);198 Frac::int(table.get(&full, which))199 })200 .collect()201}202203pub fn build(which: usize, table: &mut Table) -> Rep {204 let suffix = suffixes();205 let mut space = Space::new();206 let mut basis: Vec<Vec<u8>> = Vec::new();207 let mut queue: Vec<Vec<u8>> = vec![Vec::new()];208 let mut head = 0usize;209 while head < queue.len() {210 let word = queue[head].clone();211 head += 1;212 let row = row_of(&word, which, &suffix, table);213 if !space.insert(row) {214 continue;215 }216 basis.push(word.clone());217 for &code in CODES.iter() {218 let mut next = word.clone();219 next.push(code);220 queue.push(next);221 }222 }223 let size = basis.len();224 let mut matrices: HashMap<u8, Vec<Vec<i64>>> = HashMap::new();225 for &code in CODES.iter() {226 let mut matrix = vec![vec![0i64; size]; size];227 for (index, word) in basis.iter().enumerate() {228 let mut next = word.clone();229 next.push(code);230 let row = row_of(&next, which, &suffix, table);231 let (coefficients, rest) = space.reduce(&row);232 assert!(233 rest.iter().all(|value| value.is_zero()),234 "the basis spans every extension"235 );236 for (slot, value) in matrix[index].iter_mut().zip(coefficients.iter()) {237 *slot = value.integer();238 }239 }240 matrices.insert(code, matrix);241 }242 let empty: Vec<u8> = Vec::new();243 let (coefficients, rest) = space.reduce(&row_of(&empty, which, &suffix, table));244 assert!(245 rest.iter().all(|value| value.is_zero()),246 "the empty word is in the span"247 );248 let lambda: Vec<i64> = coefficients.iter().map(|value| value.integer()).collect();249 let gamma: Vec<i64> = basis.iter().map(|word| table.get(word, which)).collect();250 Rep {251 basis,252 matrices,253 lambda,254 gamma,255 }256}257258impl Rep {259 pub fn predict(&self, word: &[u8]) -> i64 {260 let mut state: Vec<i128> = self.lambda.iter().map(|value| *value as i128).collect();261 for code in word {262 let matrix = &self.matrices[code];263 let mut next = vec![0i128; state.len()];264 for (i, weight) in state.iter().enumerate() {265 if *weight == 0 {266 continue;267 }268 for (j, slot) in next.iter_mut().enumerate() {269 *slot += weight * matrix[i][j] as i128;270 }271 }272 state = next;273 }274 state275 .iter()276 .zip(self.gamma.iter())277 .map(|(a, b)| a * *b as i128)278 .sum::<i128>() as i64279 }280281 pub fn classes(&self) -> Vec<(Vec<u8>, Vec<Vec<i64>>)> {282 let mut out: Vec<(Vec<u8>, Vec<Vec<i64>>)> = Vec::new();283 for &code in CODES.iter() {284 let matrix = self.matrices[&code].clone();285 match out.iter_mut().find(|(_, seen)| *seen == matrix) {286 Some((members, _)) => members.push(code),287 None => out.push((vec![code], matrix)),288 }289 }290 out291 }292}293294pub fn product(a: &[Vec<i64>], b: &[Vec<i64>]) -> Vec<Vec<i64>> {295 let size = a.len();296 let mut out = vec![vec![0i64; size]; size];297 for i in 0..size {298 for k in 0..size {299 if a[i][k] == 0 {300 continue;301 }302 for j in 0..size {303 out[i][j] += a[i][k] * b[k][j];304 }305 }306 }307 out308}