use super::tokenizer::Tokenizer;
fn encode_representable(tokenizer: &Tokenizer, phrase: &str) -> Option<Vec<usize>> {
if let Some(ids) = tokenizer.encode_phrase(phrase) {
return Some(ids);
}
let lowercased = phrase.to_lowercase();
if lowercased != phrase
&& let Some(ids) = tokenizer.encode_phrase(&lowercased)
{
return Some(ids);
}
let folded = lowercased.replace('ё', "е");
if folded != lowercased {
return tokenizer.encode_phrase(&folded);
}
None
}
struct TrieNode {
children: std::collections::HashMap<usize, usize>,
is_end: bool,
shortest_phrase: usize,
depth: usize,
is_entry: bool,
grant: f32,
}
impl TrieNode {
fn new(depth: usize) -> Self {
Self {
children: std::collections::HashMap::new(),
is_end: false,
shortest_phrase: usize::MAX,
depth,
is_entry: false,
grant: 0.0,
}
}
}
#[derive(Clone, Copy, Default)]
pub(crate) struct BiasPath {
node: usize,
pending: f32,
}
impl BiasPath {
pub(crate) fn pending(&self) -> f32 {
self.pending
}
}
pub struct Biaser {
nodes: Vec<TrieNode>,
boost: f32,
phrase_count: usize,
}
impl Biaser {
#[cfg(test)]
pub(crate) fn from_sequences(sequences: Vec<Vec<usize>>, boost: f32) -> Option<Self> {
Self::build(sequences, boost, false)
}
fn build(sequences: Vec<Vec<usize>>, boost: f32, leading_is_entry: bool) -> Option<Self> {
let mut nodes = vec![TrieNode::new(0)];
let mut phrase_count = 0;
let entry_tokens = usize::from(leading_is_entry);
for seq in sequences {
if seq.is_empty() {
continue;
}
phrase_count += 1;
let scored = seq.len().saturating_sub(entry_tokens).max(1);
let mut node = 0usize;
for tok in seq {
node = match nodes[node].children.get(&tok) {
Some(&child) => child,
None => {
let depth = nodes[node].depth + 1;
let child = nodes.len();
nodes.push(TrieNode::new(depth));
nodes[node].children.insert(tok, child);
child
}
};
nodes[node].shortest_phrase = nodes[node].shortest_phrase.min(scored);
}
nodes[node].is_end = true;
}
if phrase_count == 0 {
return None;
}
for node in nodes.iter_mut().skip(1) {
node.is_entry = leading_is_entry && node.depth == 1;
node.grant = if node.is_entry {
0.0
} else {
boost / node.shortest_phrase.max(1) as f32
};
}
Some(Self {
nodes,
boost,
phrase_count,
})
}
pub fn from_phrases(
tokenizer: &Tokenizer,
phrases: &[(String, f32)],
boost: f32,
) -> Option<Self> {
if boost <= 0.0 {
return None;
}
let mut sequences = Vec::new();
let mut dropped: Vec<&str> = Vec::new();
for (phrase, weight) in phrases {
if *weight <= 0.0 {
continue;
}
match encode_representable(tokenizer, phrase) {
Some(ids) => sequences.push(ids),
None => dropped.push(phrase),
}
}
if !dropped.is_empty() {
tracing::warn!(
"{} hotword phrase(s) dropped, not representable in the active vocab: {}",
dropped.len(),
dropped.join(", ")
);
}
Self::build(sequences, boost, true)
}
pub fn phrase_count(&self) -> usize {
self.phrase_count
}
pub(crate) fn score_token(&self, path: BiasPath, tok: usize) -> (f32, BiasPath) {
let enter = |from: BiasPath, child: usize| {
let share = self.nodes[child].grant;
(
share,
BiasPath {
node: child,
pending: if self.nodes[child].is_end {
0.0
} else {
from.pending + share
},
},
)
};
if let Some(&child) = self.nodes[path.node].children.get(&tok) {
return enter(path, child);
}
let refund = -path.pending;
match self.nodes[0].children.get(&tok) {
Some(&child) => {
let (delta, next) = enter(BiasPath::default(), child);
(refund + delta, next)
}
None => (refund, BiasPath::default()),
}
}
pub(crate) fn continuations(&self, path: BiasPath, out: &mut Vec<usize>) {
out.extend(self.nodes[path.node].children.keys().copied());
if path.node != 0 {
out.extend(self.nodes[0].children.keys().copied());
}
}
pub(crate) fn new_state(&self) -> BiasState {
BiasState {
active: vec![0],
}
}
pub(crate) fn boost_logits(&self, state: &BiasState, logits: &mut [f32]) {
for &node in &state.active {
for (&tok, &child) in &self.nodes[node].children {
if tok < logits.len() && !self.nodes[child].is_entry {
logits[tok] += self.boost;
}
}
}
}
pub(crate) fn advance(&self, state: &mut BiasState, tok: usize) {
let mut next = Vec::new();
for &node in &state.active {
if let Some(&child) = self.nodes[node].children.get(&tok)
&& !next.contains(&child)
{
next.push(child);
}
}
if !next.contains(&0) {
next.push(0);
}
state.active = next;
}
}
pub(crate) struct BiasState {
active: Vec<usize>,
}
#[cfg(test)]
mod tests {
use super::*;
fn biaser(seqs: Vec<Vec<usize>>, boost: f32) -> Biaser {
Biaser::from_sequences(seqs, boost).expect("non-empty sequences")
}
#[test]
fn test_from_sequences_empty_returns_none() {
assert!(Biaser::from_sequences(vec![], 5.0).is_none());
assert!(Biaser::from_sequences(vec![vec![]], 5.0).is_none());
}
#[test]
fn test_boost_applies_to_first_token_of_each_hotword() {
let b = biaser(vec![vec![1, 2], vec![3]], 5.0);
let state = b.new_state();
let mut logits = vec![0.0; 5];
b.boost_logits(&state, &mut logits);
assert_eq!(logits[1], 5.0);
assert_eq!(logits[3], 5.0);
assert_eq!(logits[2], 0.0, "mid-hotword token not boosted at root");
assert_eq!(logits[0], 0.0);
}
#[test]
fn test_advance_then_boost_continuation() {
let b = biaser(vec![vec![1, 2]], 5.0);
let mut state = b.new_state();
b.advance(&mut state, 1);
let mut logits = vec![0.0; 5];
b.boost_logits(&state, &mut logits);
assert_eq!(logits[2], 5.0, "continuation token boosted after prefix");
assert_eq!(logits[1], 5.0, "root keeps a fresh hotword start available");
}
#[test]
fn test_advance_off_prefix_resets_to_root_only() {
let b = biaser(vec![vec![1, 2]], 5.0);
let mut state = b.new_state();
b.advance(&mut state, 1); b.advance(&mut state, 9); let mut logits = vec![0.0; 5];
b.boost_logits(&state, &mut logits);
assert_eq!(logits[2], 0.0, "continuation no longer boosted after reset");
assert_eq!(logits[1], 5.0, "root start still boosted");
}
#[test]
fn test_shared_prefix_keeps_both_branches_active() {
let b = biaser(vec![vec![1, 2], vec![1, 3]], 4.0);
let mut state = b.new_state();
b.advance(&mut state, 1);
let mut logits = vec![0.0; 5];
b.boost_logits(&state, &mut logits);
assert_eq!(logits[2], 4.0);
assert_eq!(logits[3], 4.0);
}
#[test]
fn test_boost_ignores_out_of_range_token_id() {
let b = biaser(vec![vec![99]], 5.0);
let state = b.new_state();
let mut logits = vec![0.0; 5];
b.boost_logits(&state, &mut logits); assert!(logits.iter().all(|&l| l == 0.0));
}
use crate::inference::tokenizer::Tokenizer;
fn char_tokenizer() -> Tokenizer {
let tokens = vec![
"а".to_string(),
"б".to_string(),
"в".to_string(),
"г".to_string(),
"д".to_string(),
"\u{2581}".to_string(), "<unk>".to_string(),
"<blk>".to_string(),
];
Tokenizer::from_tokens(tokens)
}
#[test]
fn the_phrase_entry_marker_is_never_boosted() {
let tok = char_tokenizer();
let marker = tok.encode_phrase("а").expect("representable")[0];
for phrases in [
vec![("а".to_string(), 1.0)],
vec![("аб".to_string(), 1.0)],
vec![("вг".to_string(), 1.0), ("д".to_string(), 1.0)],
] {
let b = Biaser::from_phrases(&tok, &phrases, 6.0).expect("compiles");
let state = b.new_state();
let mut logits = vec![0.0; 8];
b.boost_logits(&state, &mut logits);
assert_eq!(
logits[marker], 0.0,
"the entry marker took a boost for {phrases:?}"
);
assert!(
logits.iter().all(|&l| l == 0.0),
"nothing is boostable before a word boundary is emitted"
);
}
}
#[test]
fn a_phrase_earns_one_boost_however_long_it_is() {
for len in [1usize, 3, 9] {
let seq: Vec<usize> = (1..=len).collect();
let b = Biaser::from_sequences(vec![seq.clone()], 6.0).expect("biaser");
let mut path = BiasPath::default();
let mut earned = 0.0;
for &tok in &seq {
let (delta, next) = b.score_token(path, tok);
earned += delta;
path = next;
}
assert!(
(earned - 6.0).abs() < 1e-4,
"a {len}-token phrase earned {earned}, expected the boost once"
);
assert_eq!(path.pending(), 0.0, "a finished phrase owes nothing back");
}
}
#[test]
fn an_abandoned_phrase_earns_nothing() {
let b = Biaser::from_sequences(vec![vec![1, 2, 3, 4]], 6.0).expect("biaser");
let mut path = BiasPath::default();
let mut earned = 0.0;
for tok in [1usize, 2] {
let (delta, next) = b.score_token(path, tok);
earned += delta;
path = next;
}
assert!(path.pending() > 0.0, "mid-phrase credit is outstanding");
let (delta, path) = b.score_token(path, 99);
earned += delta;
assert!(earned.abs() < 1e-4, "abandoned phrase netted {earned}");
assert_eq!(path.pending(), 0.0);
}
#[test]
fn test_from_phrases_zero_boost_returns_none() {
let tok = char_tokenizer();
let phrases = vec![("аб".to_string(), 1.0)];
assert!(Biaser::from_phrases(&tok, &phrases, 0.0).is_none());
}
#[test]
fn test_from_phrases_negative_boost_returns_none() {
let tok = char_tokenizer();
let phrases = vec![("аб".to_string(), 1.0)];
assert!(Biaser::from_phrases(&tok, &phrases, -3.0).is_none());
}
#[test]
fn test_from_phrases_empty_slice_returns_none() {
let tok = char_tokenizer();
assert!(Biaser::from_phrases(&tok, &[], 5.0).is_none());
}
#[test]
fn test_from_phrases_all_zero_weight_returns_none() {
let tok = char_tokenizer();
let phrases = vec![("аб".to_string(), 0.0), ("вг".to_string(), -1.0)];
assert!(Biaser::from_phrases(&tok, &phrases, 5.0).is_none());
}
#[test]
fn test_from_phrases_unrepresentable_only_returns_none() {
let tok = char_tokenizer();
let phrases = vec![("xyz".to_string(), 1.0)];
assert!(Biaser::from_phrases(&tok, &phrases, 5.0).is_none());
}
#[test]
fn test_from_phrases_single_token_phrase_boosts_first_token() {
let tok = char_tokenizer();
let phrases = vec![("а".to_string(), 1.0)];
let b = Biaser::from_phrases(&tok, &phrases, 7.0).expect("phrase compiles");
assert_eq!(b.phrase_count(), 1);
let ids = tok.encode_phrase("а").expect("representable");
let mut state = b.new_state();
let mut logits = vec![0.0; 8];
b.boost_logits(&state, &mut logits);
assert_eq!(logits[ids[0]], 0.0, "the boundary marker is never boosted");
b.advance(&mut state, ids[0]);
let mut logits = vec![0.0; 8];
b.boost_logits(&state, &mut logits);
assert_eq!(
logits[ids[1]], 7.0,
"a one-character phrase is worth it all"
);
}
#[test]
fn test_from_phrases_multi_token_phrase_boosts_continuation() {
let tok = char_tokenizer();
let phrases = vec![("аб".to_string(), 1.0)];
let b = Biaser::from_phrases(&tok, &phrases, 4.0).expect("phrase compiles");
assert_eq!(b.phrase_count(), 1);
let ids = tok.encode_phrase("аб").expect("representable");
assert_eq!(ids, vec![5, 0, 1]);
let mut state = b.new_state();
b.advance(&mut state, ids[0]); b.advance(&mut state, ids[1]); let mut logits = vec![0.0; 8];
b.boost_logits(&state, &mut logits);
assert_eq!(
logits[ids[2]], 4.0,
"third token boosted after two-token prefix"
);
}
#[test]
fn test_from_phrases_drops_unrepresentable_keeps_representable() {
let tok = char_tokenizer();
let phrases = vec![("аб".to_string(), 1.0), ("аz".to_string(), 1.0)];
let b = Biaser::from_phrases(&tok, &phrases, 5.0).expect("one phrase compiles");
assert_eq!(b.phrase_count(), 1);
}
fn cyrillic_tokenizer() -> Tokenizer {
let tokens = [
"г", "и", "а", "э", "м", "п", "т", "р", "е", "\u{2581}", "<unk>", "<blk>",
];
Tokenizer::from_tokens(tokens.iter().map(|t| (*t).to_string()).collect())
}
#[test]
fn test_from_phrases_capitalized_phrase_is_kept() {
let tok = cyrillic_tokenizer();
let phrases = vec![("Гигаэм".to_string(), 1.0)];
let b = Biaser::from_phrases(&tok, &phrases, 5.0).expect("capitalized phrase compiles");
assert_eq!(b.phrase_count(), 1);
let ids = tok.encode_phrase("гигаэм").expect("representable");
let mut state = b.new_state();
b.advance(&mut state, ids[0]); let mut logits = vec![0.0; 12];
b.boost_logits(&state, &mut logits);
assert!(
logits[ids[1]] > 0.0,
"the lowercased spelling is what biases"
);
}
#[test]
fn test_from_phrases_yo_folds_to_e_when_the_vocab_lacks_it() {
let tok = cyrillic_tokenizer();
assert!(tok.encode_phrase("пётр").is_none(), "vocab has no ё");
let phrases = vec![("Пётр".to_string(), 1.0)];
let b = Biaser::from_phrases(&tok, &phrases, 5.0).expect("ё folds to е");
assert_eq!(b.phrase_count(), 1);
}
#[test]
fn test_from_phrases_cased_vocab_keeps_the_written_spelling() {
let tokens = ["Аб", "а", "б", "\u{2581}", "<unk>", "<blk>"];
let tok = Tokenizer::from_tokens(tokens.iter().map(|t| (*t).to_string()).collect());
let written = tok.encode_phrase("Аб").expect("cased vocab represents it");
let b = Biaser::from_phrases(&tok, &[("Аб".to_string(), 1.0)], 5.0).expect("compiles");
let mut state = b.new_state();
b.advance(&mut state, written[0]); let mut logits = vec![0.0; 6];
b.boost_logits(&state, &mut logits);
assert_eq!(
logits[written[1]], 5.0,
"the written form must survive untouched"
);
}
#[test]
fn test_from_phrases_latin_stays_unrepresentable_on_a_cyrillic_vocab() {
let tok = cyrillic_tokenizer();
let phrases = vec![("ChatGPT".to_string(), 1.0), ("Гигаэм".to_string(), 1.0)];
let b = Biaser::from_phrases(&tok, &phrases, 5.0).expect("one phrase compiles");
assert_eq!(b.phrase_count(), 1);
}
#[test]
fn test_from_phrases_weight_filters_per_phrase() {
let tok = char_tokenizer();
let phrases = vec![("аб".to_string(), 1.0), ("вг".to_string(), 0.0)];
let b = Biaser::from_phrases(&tok, &phrases, 5.0).expect("one phrase compiles");
assert_eq!(b.phrase_count(), 1);
}
}