Download rust/src/engine.rs from TheREZOR/TinyDecide: direct link, hf CLI and curl.
- Browser
- Download file 34.3 kB
-
https://huggingface.co/TheREZOR/TinyDecide/resolve/main/rust/src/engine.rs
- Command line
-
hf download hf://TheREZOR/TinyDecide/rust/src/engine.rs
-
curl -L -o engine.rs https://huggingface.co/TheREZOR/TinyDecide/resolve/main/rust/src/engine.rs
34.3 kB
| //! The encoder and the typed heads. A line-by-line port of `TinyDecide` in tinydecide.js. | |
| use std::path::Path; | |
| use std::time::Instant; | |
| use serde::{Deserialize, Serialize}; | |
| use crate::error::{Error, Result}; | |
| use crate::tokenizer::{is_js_space, WordPiece}; | |
| use crate::weights::{Meta, Tensor, Weights}; | |
| const K_STATE: u8 = 0; | |
| const K_QTEXT: u8 = 1; | |
| const K_ANS: u8 = 2; | |
| const K_OPT: u8 = 3; | |
| const K_LV: u8 = 4; | |
| const K_BUCKETS: [i32; 4] = [1, 2, 4, 8]; | |
| pub(crate) fn bucket(c: i32) -> usize { | |
| K_BUCKETS.iter().filter(|&&e| c > e).count() | |
| } | |
| /// The four answer types. | |
| pub enum QType { | |
| /// Pick one of 2 to 32 options. | |
| Choice, | |
| /// Is a statement about the message true? | |
| Noul, | |
| /// Place the message on ordered levels, lowest first. | |
| Score, | |
| /// Extract a piece of the message. | |
| Span, | |
| } | |
| impl QType { | |
| pub fn as_str(self) -> &'static str { | |
| match self { | |
| QType::Choice => "choice", | |
| QType::Noul => "noul", | |
| QType::Score => "score", | |
| QType::Span => "span", | |
| } | |
| } | |
| } | |
| /// One question, written in plain language at call time. | |
| pub struct Question { | |
| pub kind: QType, | |
| pub text: String, | |
| /// Options (choice) or levels, lowest first (score). Empty for noul and span. | |
| pub options: Vec<String>, | |
| } | |
| impl Question { | |
| pub fn choice<S: AsRef<str>>(text: &str, options: &[S]) -> Self { | |
| Question { kind: QType::Choice, text: text.into(), options: options.iter().map(|o| o.as_ref().to_string()).collect() } | |
| } | |
| pub fn noul(text: &str) -> Self { | |
| Question { kind: QType::Noul, text: text.into(), options: vec![] } | |
| } | |
| pub fn score<S: AsRef<str>>(text: &str, levels: &[S]) -> Self { | |
| Question { kind: QType::Score, text: text.into(), options: levels.iter().map(|o| o.as_ref().to_string()).collect() } | |
| } | |
| pub fn span(text: &str) -> Self { | |
| Question { kind: QType::Span, text: text.into(), options: vec![] } | |
| } | |
| } | |
| /// Corrections for one question (see [`crate::corrections::make_protos`]). | |
| pub struct Protos { | |
| /// One prototype per option (2 for noul: false, true), `n_options * qvec_len` values. | |
| pub vec: Vec<f32>, | |
| /// Examples behind each prototype; options with 0 add nothing. | |
| pub cnt: Vec<i32>, | |
| /// Mean vector the cosines are centred on. Without it the engine falls back to the older, | |
| /// uncentred rule, as tinydecide.js does. | |
| pub center: Option<Vec<f32>>, | |
| /// Trust weight. `None` means 1. | |
| pub lam: Option<f64>, | |
| } | |
| /// One answer. Which fields are set depends on `kind`, as in the JavaScript result. | |
| pub struct Answer { | |
| pub kind: QType, | |
| /// choice / score: one probability per option. | |
| pub probs: Vec<f64>, | |
| /// choice / score: index of the most likely option. | |
| pub pick: Option<usize>, | |
| /// choice / score: 1 - normalised entropy. | |
| pub confidence: Option<f64>, | |
| /// score: expected level from 0 (first) to 1 (last). | |
| pub score: Option<f64>, | |
| /// noul: probability the statement is true. | |
| pub p: Option<f64>, | |
| /// choice / score / noul: the vector to store with a correction. | |
| pub qvec: Vec<f32>, | |
| /// True when `qvec` is the prototype projection (every released model). | |
| pub proj: bool, | |
| /// choice / score / noul: zero-shot logits over temperature (noul: `[0, z]`), stored with a correction. | |
| pub z0: Vec<f64>, | |
| /// span: probability the message contains what was asked for. | |
| pub p_present: Option<f64>, | |
| /// span: probability of the chosen start and end tokens. | |
| pub p_span: Option<f64>, | |
| /// span: first and last state token of the span (0-based, without the state marker). | |
| pub tok: Option<(usize, usize)>, | |
| /// span: the extracted text, trimmed. | |
| pub text: Option<String>, | |
| /// span: `(start, end)` **UTF-8 byte offsets** into the state string (JavaScript reports UTF-16 units). | |
| pub char: Option<(usize, usize)>, | |
| } | |
| impl Answer { | |
| fn empty(kind: QType) -> Self { | |
| Answer { | |
| kind, | |
| probs: vec![], | |
| pick: None, | |
| confidence: None, | |
| score: None, | |
| p: None, | |
| qvec: vec![], | |
| proj: false, | |
| z0: vec![], | |
| p_present: None, | |
| p_span: None, | |
| tok: None, | |
| text: None, | |
| char: None, | |
| } | |
| } | |
| } | |
| pub struct Tokens { | |
| /// State tokens including the state marker. | |
| pub state: usize, | |
| pub questions: usize, | |
| pub total: usize, | |
| } | |
| pub struct Response { | |
| pub answers: Vec<Answer>, | |
| pub ids: Vec<u32>, | |
| pub tokens: Tokens, | |
| /// The message was longer than the model reads (`format.ts_max - 1` tokens) and was cut. | |
| pub truncated: bool, | |
| pub ms: f64, | |
| } | |
| // ------------------------------------------------------------------ math | |
| // 16 independent accumulators so the loop vectorises | |
| fn dot(a: &[f32], b: &[f32]) -> f32 { | |
| let n = a.len().min(b.len()); | |
| let (a, b) = (&a[..n], &b[..n]); | |
| let mut acc = [0f32; 16]; | |
| let ca = a.chunks_exact(16); | |
| let cb = b.chunks_exact(16); | |
| let (ra, rb) = (ca.remainder(), cb.remainder()); | |
| for (x, y) in ca.zip(cb) { | |
| for i in 0..16 { | |
| acc[i] += x[i] * y[i]; | |
| } | |
| } | |
| let mut s = 0f32; | |
| for v in acc { | |
| s += v; | |
| } | |
| for (x, y) in ra.iter().zip(rb) { | |
| s += x * y; | |
| } | |
| s | |
| } | |
| fn dot64(a: &[f32], b: &[f32]) -> f64 { | |
| a.iter().zip(b).map(|(&x, &y)| x as f64 * y as f64).sum() | |
| } | |
| fn dotd(a: &[f64], b: &[f64]) -> f64 { | |
| a.iter().zip(b).map(|(x, y)| x * y).sum() | |
| } | |
| fn norm64(a: &[f32]) -> Vec<f64> { | |
| let n = dot64(a, a).sqrt(); | |
| let n = if n == 0.0 || n.is_nan() { 1e-12 } else { n }; | |
| a.iter().map(|&z| z as f64 / n).collect() | |
| } | |
| struct Lin { | |
| w: Vec<f32>, | |
| b: Option<Vec<f32>>, | |
| rows: usize, | |
| cols: usize, | |
| } | |
| impl Lin { | |
| /// y[t] = W x[t] + b for T rows of x. | |
| fn apply(&self, x: &[f32], t: usize) -> Vec<f32> { | |
| let (din, dout) = (self.cols, self.rows); | |
| let mut y = vec![0f32; t * dout]; | |
| for j in 0..dout { | |
| let row = &self.w[j * din..(j + 1) * din]; | |
| let b = self.b.as_ref().map_or(0.0, |b| b[j]); | |
| for tt in 0..t { | |
| y[tt * dout + j] = b + dot(row, &x[tt * din..(tt + 1) * din]); | |
| } | |
| } | |
| y | |
| } | |
| fn mv(&self, v: &[f32]) -> Vec<f32> { | |
| self.apply(v, 1) | |
| } | |
| } | |
| fn layernorm(x: &[f32], t: usize, d: usize, w: &[f32], b: Option<&[f32]>, eps: f64) -> Vec<f32> { | |
| let mut y = vec![0f32; t * d]; | |
| for tt in 0..t { | |
| let r = &x[tt * d..(tt + 1) * d]; | |
| let m = r.iter().map(|&z| z as f64).sum::<f64>() / d as f64; | |
| let v = r.iter().map(|&z| (z as f64 - m) * (z as f64 - m)).sum::<f64>() / d as f64; | |
| let inv = 1.0 / (v + eps).sqrt(); | |
| for i in 0..d { | |
| let bi = b.map_or(0.0, |b| b[i] as f64); | |
| y[tt * d + i] = ((r[i] as f64 - m) * inv * w[i] as f64 + bi) as f32; | |
| } | |
| } | |
| y | |
| } | |
| /// The erf used by tinydecide.js: a Taylor series below 0.5, the Numerical Recipes erfc fit above. | |
| fn erf(x: f64) -> f64 { | |
| let s = if x > 0.0 { | |
| 1.0 | |
| } else if x < 0.0 { | |
| -1.0 | |
| } else { | |
| return x; | |
| }; | |
| let a = x.abs(); | |
| if a < 0.5 { | |
| let (mut sum, mut term, a2) = (a, a, a * a); | |
| for n in 1..12 { | |
| term *= -a2 / n as f64; | |
| sum += term / (2 * n + 1) as f64; | |
| } | |
| return s * 2.0 / std::f64::consts::PI.sqrt() * sum; | |
| } | |
| let t = 1.0 / (1.0 + 0.5 * a); | |
| let y = t | |
| * (-a * a - 1.26551223 | |
| + t * (1.00002368 | |
| + t * (0.37409196 | |
| + t * (0.09678418 | |
| + t * (-0.18628806 | |
| + t * (0.27886807 + t * (-1.13520398 + t * (1.48851587 + t * (-0.82215223 + t * 0.17087277))))))))) | |
| .exp(); | |
| s * (1.0 - y) | |
| } | |
| fn gelu(x: f32) -> f32 { | |
| let x = x as f64; | |
| (0.5 * x * (1.0 + erf(x / std::f64::consts::SQRT_2))) as f32 | |
| } | |
| /// log-softmax of a / tt. | |
| fn lsm(a: &[f64], tt: f64) -> Vec<f64> { | |
| let m = a.iter().cloned().fold(f64::NEG_INFINITY, f64::max); | |
| let l = a.iter().map(|z| ((z - m) / tt).exp()).sum::<f64>().ln(); | |
| a.iter().map(|z| (z - m) / tt - l).collect() | |
| } | |
| fn js_trim(s: &str) -> &str { | |
| s.trim_matches(is_js_space) | |
| } | |
| // ------------------------------------------------------------------ model | |
| struct Block { | |
| q: Lin, | |
| k: Lin, | |
| v: Lin, | |
| o: Lin, | |
| ln1: (Vec<f32>, Vec<f32>), | |
| fc: Lin, | |
| fc2: Lin, | |
| ln2: (Vec<f32>, Vec<f32>), | |
| } | |
| enum WordTable { | |
| Full(Vec<f32>), | |
| LowRank { a: Vec<f32>, b: Vec<f32>, r: usize }, | |
| } | |
| struct Heads { | |
| norm: (Vec<f32>, Vec<f32>), | |
| scale: Vec<f32>, | |
| a: [Lin; 2], | |
| o: [Lin; 2], | |
| p: Option<Lin>, | |
| noul_w: Vec<f32>, | |
| noul_b: f32, | |
| noul_q: Lin, | |
| sq: Lin, | |
| sk: Lin, | |
| snull_w: Vec<f32>, | |
| snull_b: f32, | |
| eq: Lin, | |
| es: Lin, | |
| ek: Lin, | |
| } | |
| struct Specials { | |
| state: u32, | |
| types: [u32; 4], | |
| sep: u32, | |
| o: u32, | |
| lv: u32, | |
| ans: u32, | |
| } | |
| struct QEnc { | |
| kind: QType, | |
| ans: usize, | |
| opt: Vec<usize>, | |
| } | |
| struct Enc { | |
| ids: Vec<u32>, | |
| pos: Vec<usize>, | |
| blk: Vec<usize>, | |
| kind: Vec<u8>, | |
| st_idx: Vec<usize>, | |
| st_off: Vec<(usize, usize)>, | |
| qs: Vec<QEnc>, | |
| truncated: bool, | |
| } | |
| /// A loaded TinyDecide model. | |
| pub struct TinyDecide { | |
| meta: Meta, | |
| tok: WordPiece, | |
| sp: Specials, | |
| word: WordTable, | |
| pos: Vec<f32>, | |
| type0: Vec<f32>, | |
| eln: (Vec<f32>, Vec<f32>), | |
| proj: Lin, | |
| blocks: Vec<Block>, | |
| heads: Heads, | |
| } | |
| fn take_vec(w: &mut Weights, n: &str, len: usize) -> Result<Vec<f32>> { | |
| Ok(w.take(n, &[Some(len)])?.data) | |
| } | |
| fn take_lin(w: &mut Weights, n: &str, rows: Option<usize>, cols: usize, bias: Option<&str>) -> Result<Lin> { | |
| let t: Tensor = w.take(n, &[rows, Some(cols)])?; | |
| let rows = t.shape[0]; | |
| let b = match bias { | |
| Some(bn) => Some(take_vec(w, bn, rows)?), | |
| None => None, | |
| }; | |
| Ok(Lin { w: t.data, b, rows, cols }) | |
| } | |
| impl TinyDecide { | |
| /// Loads `meta.json` and `model.bin` from a directory. | |
| pub fn load<P: AsRef<Path>>(dir: P) -> Result<Self> { | |
| let dir = dir.as_ref(); | |
| let meta = std::fs::read_to_string(dir.join("meta.json"))?; | |
| let bin = std::fs::read(dir.join("model.bin"))?; | |
| Self::from_bytes(&meta, &bin) | |
| } | |
| /// Builds a model from the text of `meta.json` and the bytes of `model.bin`. | |
| pub fn from_bytes(meta_json: &str, bin: &[u8]) -> Result<Self> { | |
| let mut meta: Meta = serde_json::from_str(meta_json)?; | |
| let c = meta.cfg.clone(); | |
| if c.arch != "electra" { | |
| return Err(Error::Model(format!("unsupported architecture {}", c.arch))); | |
| } | |
| if meta.tokenizer.kind != "wordpiece" { | |
| return Err(Error::Model("this engine build ships the WordPiece tokenizer only".into())); | |
| } | |
| if c.heads == 0 || c.d % c.heads != 0 { | |
| return Err(Error::Model("cfg.d must be a multiple of cfg.heads".into())); | |
| } | |
| if meta.temp.len() < 4 || meta.beta.len() < 5 { | |
| return Err(Error::Model("meta.temp needs 4 values and meta.beta 5".into())); | |
| } | |
| let tk = &meta.tokenizer; | |
| let tok = WordPiece::new(&tk.vocab, &tk.unk, &tk.prefix, tk.max_chars)?; | |
| let vocab_len = tk.vocab.len(); | |
| meta.tokenizer.vocab = Vec::new(); | |
| let mut w = Weights::new(&meta, bin)?; | |
| let (d, e, dh) = (c.d, c.emb_dim, c.dh_head); | |
| let word = if w.has("m.word.a.weight") { | |
| let a = w.take("m.word.a.weight", &[None, None])?; | |
| let r = a.shape[1]; | |
| let b = w.take("m.word.b.weight", &[Some(e), Some(r)])?; | |
| (WordTable::LowRank { a: a.data, b: b.data, r }, a.shape[0]) | |
| } else { | |
| let t = w.take("m.word.weight", &[None, Some(e)])?; | |
| let rows = t.shape[0]; | |
| (WordTable::Full(t.data), rows) | |
| }; | |
| let (word, vocab_rows) = word; | |
| let posemb = w.take("m.posemb.weight", &[None, Some(e)])?; | |
| let pos_rows = posemb.shape[0]; | |
| let f = &meta.format; | |
| if f.ts_max == 0 || f.ts_max > pos_rows || f.p_q + f.q_max > pos_rows { | |
| return Err(Error::Model("format.ts_max / p_q / q_max do not fit the position table".into())); | |
| } | |
| let spc = |n: &str| -> Result<u32> { | |
| let id = *meta.specials.get(n).ok_or_else(|| Error::Model(format!("missing special token {n}")))?; | |
| if id as usize >= vocab_rows { | |
| return Err(Error::Model(format!("special token {n} is outside the embedding table"))); | |
| } | |
| Ok(id) | |
| }; | |
| let sp = Specials { | |
| state: spc("<|state|>")?, | |
| types: [spc("<|choice|>")?, spc("<|noul|>")?, spc("<|score|>")?, spc("<|span|>")?], | |
| sep: spc("<|sep|>")?, | |
| o: spc("<|o|>")?, | |
| lv: spc("<|lv|>")?, | |
| ans: spc("<|ans|>")?, | |
| }; | |
| if tok.unk() as usize >= vocab_rows { | |
| return Err(Error::Model("the unknown token is outside the embedding table".into())); | |
| } | |
| // every vocabulary id must have an embedding row | |
| if vocab_len > vocab_rows { | |
| return Err(Error::Model("the vocabulary is larger than the embedding table".into())); | |
| } | |
| let type0 = take_vec(&mut w, "m.type0", e)?; | |
| let eln = (take_vec(&mut w, "m.eln.weight", e)?, take_vec(&mut w, "m.eln.bias", e)?); | |
| let proj = take_lin(&mut w, "m.proj.weight", Some(d), e, Some("m.proj.bias"))?; | |
| let mut blocks = Vec::with_capacity(c.layers); | |
| for l in 0..c.layers { | |
| let p = format!("m.blocks.{l}."); | |
| let n = |s: &str| format!("{p}{s}"); | |
| let q = take_lin(&mut w, &n("q.weight"), Some(d), d, Some(&n("q.bias")))?; | |
| let k = take_lin(&mut w, &n("k.weight"), Some(d), d, Some(&n("k.bias")))?; | |
| let v = take_lin(&mut w, &n("v.weight"), Some(d), d, Some(&n("v.bias")))?; | |
| let o = take_lin(&mut w, &n("o.weight"), Some(d), d, Some(&n("o.bias")))?; | |
| let ln1 = (take_vec(&mut w, &n("ln1.weight"), d)?, take_vec(&mut w, &n("ln1.bias"), d)?); | |
| let fc = take_lin(&mut w, &n("fc.weight"), None, d, Some(&n("fc.bias")))?; | |
| let fc2 = take_lin(&mut w, &n("fc2.weight"), Some(d), fc.rows, Some(&n("fc2.bias")))?; | |
| let ln2 = (take_vec(&mut w, &n("ln2.weight"), d)?, take_vec(&mut w, &n("ln2.bias"), d)?); | |
| blocks.push(Block { q, k, v, o, ln1, fc, fc2, ln2 }); | |
| } | |
| let hl = |w: &mut Weights, n: &str| take_lin(w, n, Some(dh), d, None); | |
| let p = if w.has("h.P.weight") { Some(hl(&mut w, "h.P.weight")?) } else { None }; | |
| let heads = Heads { | |
| norm: (take_vec(&mut w, "h.norm.weight", d)?, take_vec(&mut w, "h.norm.bias", d)?), | |
| scale: take_vec(&mut w, "h.scale", 2)?, | |
| a: [hl(&mut w, "h.A.0.weight")?, hl(&mut w, "h.A.1.weight")?], | |
| o: [hl(&mut w, "h.O.0.weight")?, hl(&mut w, "h.O.1.weight")?], | |
| p, | |
| noul_w: w.take("h.noul.weight", &[Some(1), Some(d)])?.data, | |
| noul_b: take_vec(&mut w, "h.noul.bias", 1)?[0], | |
| noul_q: hl(&mut w, "h.noul_q.weight")?, | |
| sq: hl(&mut w, "h.sq.weight")?, | |
| sk: hl(&mut w, "h.sk.weight")?, | |
| snull_w: w.take("h.snull.weight", &[Some(1), Some(d)])?.data, | |
| snull_b: take_vec(&mut w, "h.snull.bias", 1)?[0], | |
| eq: hl(&mut w, "h.eq.weight")?, | |
| es: hl(&mut w, "h.es.weight")?, | |
| ek: hl(&mut w, "h.ek.weight")?, | |
| }; | |
| Ok(TinyDecide { meta, tok, sp, word, pos: posemb.data, type0, eln, proj, blocks, heads }) | |
| } | |
| /// The model's metadata (`cfg`, `format`, `temp`, `beta`, ...); the vocabulary is dropped after loading. | |
| pub fn meta(&self) -> &Meta { | |
| &self.meta | |
| } | |
| /// `meta.beta`, the prototype weights by example count; pass it to [`crate::corrections::make_protos`]. | |
| pub fn beta(&self) -> &[f64] { | |
| &self.meta.beta | |
| } | |
| pub fn tokenizer(&self) -> &WordPiece { | |
| &self.tok | |
| } | |
| /// Answers every question about `state` in one encoder pass. | |
| pub fn answer(&self, state: &str, questions: &[Question]) -> Result<Response> { | |
| self.answer_with(state, questions, &[]) | |
| } | |
| /// Like [`answer`](Self::answer), with corrections: `protos[i]` belongs to `questions[i]` | |
| /// (missing entries mean none). | |
| pub fn answer_with(&self, state: &str, questions: &[Question], protos: &[Option<Protos>]) -> Result<Response> { | |
| let t0 = Instant::now(); | |
| let enc = self.encode_request(state, questions)?; | |
| let dh = self.meta.cfg.dh_head; | |
| for (qi, q) in questions.iter().enumerate() { | |
| if let Some(Some(p)) = protos.get(qi) { | |
| let k = if q.kind == QType::Noul { 2 } else { q.options.len() }; | |
| for i in 0..k { | |
| if p.cnt.get(i).copied().unwrap_or(0) > 0 && p.vec.len() < (i + 1) * dh { | |
| return Err(Error::Request(format!("Corrections for question {} have too few vector values.", qi + 1))); | |
| } | |
| } | |
| if p.center.as_ref().is_some_and(|c| c.len() < dh) { | |
| return Err(Error::Request(format!("Corrections for question {} have a short center.", qi + 1))); | |
| } | |
| } | |
| } | |
| let x = self.hidden(&enc); | |
| let answers = self.heads(&enc, &x, state, protos); | |
| let t = enc.ids.len(); | |
| let s = enc.st_idx.len(); | |
| Ok(Response { | |
| answers, | |
| tokens: Tokens { state: s + 1, questions: t - s - 1, total: t }, | |
| truncated: enc.truncated, | |
| ms: t0.elapsed().as_secs_f64() * 1000.0, | |
| ids: enc.ids, | |
| }) | |
| } | |
| fn encode_request(&self, state: &str, questions: &[Question]) -> Result<Enc> { | |
| let f = &self.meta.format; | |
| let st = self.tok.encode(state); | |
| let n = st.ids.len().min(f.ts_max - 1); | |
| let mut ids = Vec::with_capacity(n + 1 + questions.len() * 16); | |
| ids.push(self.sp.state); | |
| ids.extend_from_slice(&st.ids[..n]); | |
| let mut pos: Vec<usize> = (0..ids.len()).collect(); | |
| let mut blk = vec![0usize; ids.len()]; | |
| let mut kind = vec![K_STATE; ids.len()]; | |
| let st_idx: Vec<usize> = (1..ids.len()).collect(); | |
| let st_off = st.offsets[..n].to_vec(); | |
| let mut qs = Vec::with_capacity(questions.len()); | |
| for (qi, q) in questions.iter().enumerate() { | |
| let k = qi + 1; | |
| let k_max = f.k_max.unwrap_or(usize::MAX); | |
| let n_opt = q.options.len(); | |
| if matches!(q.kind, QType::Choice | QType::Score) && (n_opt < 2 || n_opt > k_max) { | |
| let max = f.k_max.map_or("any".to_string(), |m| m.to_string()); | |
| return Err(Error::Request(format!("Question {k} needs 2 to {max} options, not {n_opt}."))); | |
| } | |
| let ti = match q.kind { | |
| QType::Choice => 0, | |
| QType::Noul => 1, | |
| QType::Score => 2, | |
| QType::Span => 3, | |
| }; | |
| let mut b_ids = vec![self.sp.types[ti]]; | |
| let mut b_kind = vec![K_QTEXT]; | |
| let t = self.tok.encode(&q.text).ids; | |
| b_kind.extend(std::iter::repeat_n(K_QTEXT, t.len())); | |
| b_ids.extend(t); | |
| let mut opt_local = Vec::new(); | |
| if !q.options.is_empty() { | |
| b_ids.push(self.sp.sep); | |
| b_kind.push(K_QTEXT); | |
| let (mk, mkind) = if q.kind == QType::Choice { (self.sp.o, K_OPT) } else { (self.sp.lv, K_LV) }; | |
| for o in &q.options { | |
| let oi = self.tok.encode(o).ids; | |
| b_kind.extend(std::iter::repeat_n(K_QTEXT, oi.len())); | |
| b_ids.extend(oi); | |
| opt_local.push(b_ids.len()); | |
| b_ids.push(mk); | |
| b_kind.push(mkind); | |
| } | |
| } | |
| let ans_local = b_ids.len(); | |
| b_ids.push(self.sp.ans); | |
| b_kind.push(K_ANS); | |
| if b_ids.len() > f.q_max { | |
| return Err(Error::Request(format!("Question {k} is too long ({} tokens, max {}).", b_ids.len(), f.q_max))); | |
| } | |
| let base = ids.len(); | |
| for (j, (&id, &kd)) in b_ids.iter().zip(&b_kind).enumerate() { | |
| ids.push(id); | |
| pos.push(f.p_q + j); | |
| blk.push(k); | |
| kind.push(kd); | |
| } | |
| qs.push(QEnc { kind: q.kind, ans: base + ans_local, opt: opt_local.iter().map(|j| base + j).collect() }); | |
| } | |
| Ok(Enc { ids, pos, blk, kind, st_idx, st_off, qs, truncated: st.ids.len() > n }) | |
| } | |
| fn continues(&self, kind: &[u8]) -> Vec<bool> { | |
| match self.meta.cfg.fusion.as_deref() { | |
| Some("all") => vec![true; kind.len()], | |
| Some("markers") => kind.iter().map(|&k| matches!(k, K_STATE | K_ANS | K_OPT | K_LV)).collect(), | |
| _ => kind.iter().map(|&k| k == K_STATE || k == K_ANS).collect(), | |
| } | |
| } | |
| fn allowed_sets(enc: &Enc, cont: &[bool], fusion_layer: bool) -> Vec<Vec<usize>> { | |
| let t = enc.ids.len(); | |
| let nb = enc.blk.iter().max().map_or(0, |m| m + 1); | |
| let mut by_blk: Vec<Vec<usize>> = vec![Vec::new(); nb]; | |
| for j in 0..t { | |
| if fusion_layer && !cont[j] { | |
| continue; | |
| } | |
| by_blk[enc.blk[j]].push(j); | |
| } | |
| let mut sets = Vec::with_capacity(t); | |
| for i in 0..t { | |
| if fusion_layer && !cont[i] { | |
| sets.push(vec![i]); | |
| continue; | |
| } | |
| let own = if by_blk[enc.blk[i]].is_empty() { vec![i] } else { by_blk[enc.blk[i]].clone() }; | |
| if !fusion_layer || enc.blk[i] == 0 { | |
| sets.push(own); | |
| } else { | |
| let mut s = by_blk[0].clone(); | |
| s.extend(own); | |
| sets.push(s); | |
| } | |
| } | |
| sets | |
| } | |
| fn attend(&self, q: &[f32], k: &[f32], v: &[f32], t: usize, sets: &[Vec<usize>]) -> Vec<f32> { | |
| let d = self.meta.cfg.d; | |
| let heads = self.meta.cfg.heads; | |
| let dh = d / heads; | |
| let scale = 1.0 / (dh as f32).sqrt(); | |
| let mut out = vec![0f32; t * d]; | |
| let mut sc = vec![0f64; t]; | |
| for (i, keys) in sets.iter().enumerate().take(t) { | |
| for h in 0..heads { | |
| let qo = i * d + h * dh; | |
| let qr = &q[qo..qo + dh]; | |
| let mut mx = f64::NEG_INFINITY; | |
| for (a, &kj) in keys.iter().enumerate() { | |
| let ko = kj * d + h * dh; | |
| let s = (dot(qr, &k[ko..ko + dh]) * scale) as f64; | |
| sc[a] = s; | |
| if s > mx { | |
| mx = s; | |
| } | |
| } | |
| let mut z = 0f64; | |
| for s in sc.iter_mut().take(keys.len()) { | |
| *s = (*s - mx).exp(); | |
| z += *s; | |
| } | |
| let o = &mut out[qo..qo + dh]; | |
| for (a, &kj) in keys.iter().enumerate() { | |
| let p = (sc[a] / z) as f32; | |
| let vo = kj * d + h * dh; | |
| for (oc, &vc) in o.iter_mut().zip(&v[vo..vo + dh]) { | |
| *oc += p * vc; | |
| } | |
| } | |
| } | |
| } | |
| out | |
| } | |
| fn hidden(&self, enc: &Enc) -> Vec<f32> { | |
| let c = &self.meta.cfg; | |
| let (t, d, e) = (enc.ids.len(), c.d, c.emb_dim); | |
| let l_a = c.l_a.unwrap_or(0); | |
| let cont = self.continues(&enc.kind); | |
| let sets_f = Self::allowed_sets(enc, &cont, true); | |
| let sets_c = if l_a > 0 { Self::allowed_sets(enc, &cont, false) } else { Vec::new() }; | |
| let mut raw = vec![0f32; t * e]; | |
| for tt in 0..t { | |
| let id = enc.ids[tt] as usize; | |
| let po = enc.pos[tt] * e; | |
| for ch in 0..e { | |
| let wv = match &self.word { | |
| WordTable::Full(w) => w[id * e + ch] as f64, | |
| WordTable::LowRank { a, b, r } => dot64(&a[id * r..(id + 1) * r], &b[ch * r..(ch + 1) * r]), | |
| }; | |
| raw[tt * e + ch] = (wv + self.pos[po + ch] as f64 + self.type0[ch] as f64) as f32; | |
| } | |
| } | |
| let n = layernorm(&raw, t, e, &self.eln.0, Some(&self.eln.1), c.ln_eps); | |
| let mut x = self.proj.apply(&n, t); | |
| for (l, b) in self.blocks.iter().enumerate() { | |
| let sets = if l < l_a { &sets_c } else { &sets_f }; | |
| let q = b.q.apply(&x, t); | |
| let k = b.k.apply(&x, t); | |
| let v = b.v.apply(&x, t); | |
| let y = self.attend(&q, &k, &v, t, sets); | |
| let mut o = b.o.apply(&y, t); | |
| for (oi, xi) in o.iter_mut().zip(&x) { | |
| *oi += xi; | |
| } | |
| x = layernorm(&o, t, d, &b.ln1.0, Some(&b.ln1.1), c.ln_eps); | |
| let mut f = b.fc.apply(&x, t); | |
| for z in f.iter_mut() { | |
| *z = gelu(*z); | |
| } | |
| let mut f2 = b.fc2.apply(&f, t); | |
| for (fi, xi) in f2.iter_mut().zip(&x) { | |
| *fi += xi; | |
| } | |
| x = layernorm(&f2, t, d, &b.ln2.0, Some(&b.ln2.1), c.ln_eps); | |
| } | |
| x | |
| } | |
| fn heads(&self, enc: &Enc, x: &[f32], state: &str, protos: &[Option<Protos>]) -> Vec<Answer> { | |
| let c = &self.meta.cfg; | |
| let (d, dh, t) = (c.d, c.dh_head, enc.ids.len()); | |
| let h = &self.heads; | |
| let hq = layernorm(x, t, d, &h.norm.0, Some(&h.norm.1), 1e-5); | |
| let row = |i: usize| &hq[i * d..(i + 1) * d]; | |
| let temp = &self.meta.temp; | |
| let beta = &self.meta.beta; | |
| let sqrt_dh = (dh as f64).sqrt(); | |
| let has_p = h.p.is_some(); | |
| let mut out = Vec::with_capacity(enc.qs.len()); | |
| for (qi, q) in enc.qs.iter().enumerate() { | |
| let h_ans = row(q.ans); | |
| let pr = protos.get(qi).and_then(|p| p.as_ref()); | |
| let pv: Option<Vec<f32>> = h.p.as_ref().map(|p| p.mv(h_ans)); | |
| let cnt = |i: usize| pr.and_then(|p| p.cnt.get(i).copied()).unwrap_or(0); | |
| let pvec = |i: usize| &pr.expect("checked").vec[i * dh..(i + 1) * dh]; | |
| // centred-cosine prototype term: beta(k) * cos(p - c, mean_k - c); None = the older rule | |
| let proto_term = |i: usize| -> Option<f64> { | |
| let Some(p) = pr else { return Some(0.0) }; | |
| if cnt(i) <= 0 { | |
| return Some(0.0); | |
| } | |
| match (&pv, &p.center) { | |
| (Some(pv), Some(cn)) => { | |
| let v = pvec(i); | |
| let a: Vec<f32> = (0..dh).map(|j| pv[j] - cn[j]).collect(); | |
| let b: Vec<f32> = (0..dh).map(|j| v[j] - cn[j]).collect(); | |
| let na = dot64(&a, &a).sqrt(); | |
| let nb = dot64(&b, &b).sqrt(); | |
| let na = if na == 0.0 { 1e-12 } else { na }; | |
| let nb = if nb == 0.0 { 1e-12 } else { nb }; | |
| Some(beta[bucket(cnt(i))] * dot64(&a, &b) / (na * nb)) | |
| } | |
| _ => None, | |
| } | |
| }; | |
| let lam = pr.and_then(|p| p.lam).unwrap_or(1.0); | |
| match q.kind { | |
| QType::Choice | QType::Score => { | |
| let sel = if q.kind == QType::Score { 1 } else { 0 }; | |
| let tt = temp[if sel == 1 { 2 } else { 0 }]; | |
| let qa = h.a[sel].mv(h_ans); | |
| let s = (h.scale[sel] as f64).exp(); | |
| let qv = norm64(&qa); | |
| let mut z0 = Vec::with_capacity(q.opt.len()); | |
| let logits: Vec<f64> = q | |
| .opt | |
| .iter() | |
| .enumerate() | |
| .map(|(i, &oi)| { | |
| let ov = h.o[sel].mv(row(oi)); | |
| let mut z = s * dot64(&qa, &ov) / sqrt_dh; | |
| z0.push(z / tt); | |
| match proto_term(i) { | |
| Some(tv) => z += lam * tv, | |
| None => { | |
| if cnt(i) > 0 { | |
| z += beta[bucket(cnt(i))] * dotd(&qv, &norm64(pvec(i))); | |
| } | |
| } | |
| } | |
| z | |
| }) | |
| .collect(); | |
| let m = logits.iter().cloned().fold(f64::NEG_INFINITY, f64::max); | |
| let ex: Vec<f64> = logits.iter().map(|z| ((z - m) / tt).exp()).collect(); | |
| let zs: f64 = ex.iter().sum(); | |
| let probs: Vec<f64> = ex.iter().map(|e| e / zs).collect(); | |
| let k = probs.len(); | |
| let hh = -probs.iter().map(|&p| if p > 0.0 { p * p.ln() } else { 0.0 }).sum::<f64>(); | |
| let mut pick = 0; | |
| for (i, &p) in probs.iter().enumerate() { | |
| if p > probs[pick] { | |
| pick = i; | |
| } | |
| } | |
| let mut a = Answer::empty(q.kind); | |
| a.confidence = Some(if k > 1 { 1.0 - hh / (k as f64).ln() } else { 1.0 }); | |
| if q.kind == QType::Score { | |
| a.score = Some(if k > 1 { probs.iter().enumerate().map(|(i, p)| p * i as f64 / (k - 1) as f64).sum() } else { 0.0 }); | |
| } | |
| a.pick = Some(pick); | |
| a.probs = probs; | |
| a.qvec = match &pv { | |
| Some(pv) => pv.clone(), | |
| None => qv.iter().map(|&z| z as f32).collect(), | |
| }; | |
| a.proj = has_p; | |
| a.z0 = z0; | |
| out.push(a); | |
| } | |
| QType::Noul => { | |
| let mut z = dot64(&h.noul_w, h_ans) + h.noul_b as f64; | |
| let z0 = z / temp[1]; | |
| let nv = norm64(&h.noul_q.mv(h_ans)); | |
| if pr.is_some() { | |
| let term: Vec<f64> = (0..2) | |
| .map(|i| match proto_term(i) { | |
| Some(tv) => lam * tv, | |
| None => { | |
| if cnt(i) > 0 { | |
| beta[bucket(cnt(i))] * dotd(&nv, &norm64(pvec(i))) | |
| } else { | |
| 0.0 | |
| } | |
| } | |
| }) | |
| .collect(); | |
| z += term[1] - term[0]; | |
| } | |
| let mut a = Answer::empty(QType::Noul); | |
| a.p = Some(1.0 / (1.0 + (-z / temp[1]).exp())); | |
| a.qvec = match &pv { | |
| Some(pv) => pv.clone(), | |
| None => nv.iter().map(|&z| z as f32).collect(), | |
| }; | |
| a.proj = has_p; | |
| a.z0 = vec![0.0, z0]; | |
| out.push(a); | |
| } | |
| QType::Span => { | |
| let s = enc.st_idx.len(); | |
| let tt = temp[3]; | |
| let qs = h.sq.mv(h_ans); | |
| let mut start = Vec::with_capacity(s + 1); | |
| start.push(dot64(&h.snull_w, h_ans) + h.snull_b as f64); | |
| for &i in &enc.st_idx { | |
| start.push(dot64(&qs, &h.sk.mv(row(i))) / sqrt_dh); | |
| } | |
| let ls = lsm(&start, tt); | |
| let mut best = 0; | |
| for t in 1..s { | |
| if ls[1 + t] > ls[1 + best] { | |
| best = t; | |
| } | |
| } | |
| let eq = h.eq.mv(h_ans); | |
| let es = h.es.mv(row(enc.st_idx.get(best).copied().unwrap_or(0))); | |
| let qe: Vec<f32> = eq.iter().zip(&es).map(|(a, b)| a + b).collect(); | |
| let span_max = self.meta.format.span_max; | |
| let e: Vec<f64> = enc | |
| .st_idx | |
| .iter() | |
| .enumerate() | |
| .map(|(t, &i)| if t >= best && t < best + span_max { dot64(&qe, &h.ek.mv(row(i))) / sqrt_dh } else { -1e4 }) | |
| .collect(); | |
| let le = if s > 0 { lsm(&e, tt) } else { Vec::new() }; | |
| let mut best_e = best; | |
| for t in 0..s { | |
| if le[t] > le[best_e] { | |
| best_e = t; | |
| } | |
| } | |
| let mut a = Answer::empty(QType::Span); | |
| a.p_present = Some(1.0 - ls[0].exp()); | |
| a.tok = Some((best, best_e)); | |
| a.p_span = Some(if s > 0 { (ls[1 + best] + le[best_e]).exp() } else { 0.0 }); | |
| let mut text = String::new(); | |
| if s > 0 { | |
| if let (Some(&(a0, _)), Some(&(_, b1))) = (enc.st_off.get(best), enc.st_off.get(best_e)) { | |
| a.char = Some((a0, b1)); | |
| text = state.get(a0..b1).map(js_trim).unwrap_or("").to_string(); | |
| } | |
| } | |
| a.text = Some(text); | |
| out.push(a); | |
| } | |
| } | |
| } | |
| out | |
| } | |
| } | |