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}