diff --git a/fast/src/graph.rs b/fast/src/graph.rs index 890dcff..90c81bf 100644 --- a/fast/src/graph.rs +++ b/fast/src/graph.rs @@ -1,6 +1,6 @@ use pyo3::prelude::*; use pyo3::types::PyTuple; -use rustc_hash::FxHashMap; +use rustc_hash::{FxHashMap, FxHashSet}; use std::fmt; use std::hash::{Hash, Hasher}; use std::sync::Arc; @@ -293,6 +293,96 @@ fn seq_emit_merges(nodes: &[GraphV], emit: &mut dyn FnMut(Vec)) { } } +/// First candidate in emit/traversal order that is a member of `wanted`, +/// short-circuiting the walk. Mirrors `emit_merges` ordering exactly, so it +/// yields the same winner the reference's stable `max`/`next` tie-break would, +/// without materializing the whole candidate stream. +pub fn first_merge_in(g: &GraphV, wanted: &FxHashSet>) -> Option> { + match g { + GraphV::Node(_) => None, + GraphV::Seq(nodes) => seq_first_merge_in(nodes, wanted), + GraphV::Tree { root, children } => tree_first_merge_in(root, children, wanted), + GraphV::FullConn(nodes) => fullconn_first_merge_in(nodes, wanted), + GraphV::Unconn(subs) => subs.iter().find_map(|sg| first_merge_in(sg, wanted)), + } +} + +fn seq_first_merge_in(nodes: &[GraphV], wanted: &FxHashSet>) -> Option> { + let num_nodes = nodes.len(); + let only_minimal = ONLY_MINIMAL_MERGES.load(Ordering::Relaxed); + let max_size = MAX_MERGE_SIZE.load(Ordering::Relaxed); + + if only_minimal && max_size == 2 && nodes.iter().all(|n| n.is_node()) { + for i in 0..num_nodes.saturating_sub(1) { + let m = vec![nodes[i].clone(), nodes[i + 1].clone()]; + if wanted.contains(&m) { + return Some(m); + } + } + return None; + } + + for i in 0..num_nodes { + if !nodes[i].is_node() { + if let Some(w) = first_merge_in(&nodes[i], wanted) { + return Some(w); + } + } + if only_minimal && !nodes[i].is_node() { + continue; + } + let end = std::cmp::min(i + max_size + 1, num_nodes + 1); + for j in (i + 2)..end { + if only_minimal && !nodes[j - 1].is_node() { + break; + } + let m = if j - i == 2 { + vec![nodes[i].clone(), nodes[j - 1].clone()] + } else { + nodes[i..j].to_vec() + }; + if wanted.contains(&m) { + return Some(m); + } + } + } + None +} + +fn tree_first_merge_in(root: &GraphV, children: &[GraphV], wanted: &FxHashSet>) -> Option> { + if let Some(w) = first_merge_in(root, wanted) { + return Some(w); + } + let only_minimal = ONLY_MINIMAL_MERGES.load(Ordering::Relaxed); + if !only_minimal || (root.is_node() && children.iter().all(|c| c.is_node())) { + let mut full_merge = vec![root.clone()]; + full_merge.extend(children.iter().cloned()); + if wanted.contains(&full_merge) { + return Some(full_merge); + } + } + children.iter().find_map(|c| first_merge_in(c, wanted)) +} + +fn fullconn_first_merge_in(nodes: &[GraphV], wanted: &FxHashSet>) -> Option> { + for node in nodes { + if let Some(w) = first_merge_in(node, wanted) { + return Some(w); + } + } + for i in 0..nodes.len() { + for j in 0..nodes.len() { + if i != j { + let m = vec![nodes[i].clone(), nodes[j].clone()]; + if wanted.contains(&m) { + return Some(m); + } + } + } + } + None +} + /// Weighted recount for a top-level Seq. Within-child candidates are counted /// once per unique Arc pointer and multiplied by the number of occurrences of /// that pointer; the positional cross-boundary candidates are emitted per @@ -953,7 +1043,10 @@ fn seq_try_merge(nodes: &[GraphV], token: &GraphV, merge: &[GraphV], memo: &mut let has_complex_child = nodes.iter().any(|n| !matches!(n, GraphV::Node(_))); - let has_match = if n < m || (!has_complex_child && only_minimal) { + // n-ary (m > 2) minimal merges over all-Node runs are real candidates + // (see seq_emit_merges), so scan for them too — mirrors seq_merge, which + // has no such short-circuit. + let has_match = if n < m { false } else { let mut found = false; diff --git a/fast/src/trainer.rs b/fast/src/trainer.rs index 1686976..0b3aaa4 100644 --- a/fast/src/trainer.rs +++ b/fast/src/trainer.rs @@ -1,39 +1,106 @@ use pyo3::prelude::*; use pyo3::types::PyTuple; use rayon::prelude::*; -use rustc_hash::FxHashMap; +use rustc_hash::{FxHashMap, FxHashSet}; use std::sync::Arc; use crate::graph::{graphv_to_pyobject, pyobject_to_graphv, GraphV}; use crate::units::{apply_merge_to_cluster_cache, replace_word_cache, snapshot_word_cache}; use std::collections::HashMap; -/// Pick the merge with the highest score: (len - 1) * count. -/// Ties broken by bytes (deterministic, order-independent). -fn pick_best(map: &FxHashMap, usize>) -> Option<(Vec, usize)> { - let mut best: Option<(&Vec, usize, Vec>)> = None; - for (key, &count) in map { - let score = (key.len() - 1) * count; - if score == 0 { - continue; - } - let dominated = match &best { - Some((_, bs, _)) => score < *bs, - None => false, - }; - if dominated { +// Merge score: merging a k-tuple with `count` occurrences removes (k-1)*count nodes. +fn merge_score(nodes: &[GraphV], count: usize) -> usize { + (nodes.len() - 1) * count +} + +/// One pass over `map`: returns a representative max-score candidate, whether +/// the max score is shared by more than one candidate (a tie), and that score. +/// None when every candidate scores 0. The common no-tie path avoids allocating +/// a tie set — only ties pay for `tied_at`. +fn best_and_tie(map: &FxHashMap, usize>) -> Option<(&Vec, bool, usize)> { + let mut best_score = 0; + let mut best_key: Option<&Vec> = None; + let mut tie = false; + for (k, &c) in map { + let s = merge_score(k, c); + if s == 0 { continue; } - let bytes_key: Vec> = key.iter().map(|g| g.to_bytes()).collect(); - let replace = match &best { - None => true, - Some((_, bs, bb)) => score > *bs || (score == *bs && bytes_key > *bb), - }; - if replace { - best = Some((key, score, bytes_key)); + if s > best_score { + best_score = s; + best_key = Some(k); + tie = false; + } else if s == best_score { + tie = true; } } - best.map(|(k, _, _)| (k.clone(), map[k])) + best_key.map(|k| (k, tie, best_score)) +} + +/// Candidates whose score equals `score`. +fn tied_at(map: &FxHashMap, usize>, score: usize) -> FxHashSet> { + map.iter() + .filter(|(k, &c)| merge_score(k, c) == score) + .map(|(k, _)| k.clone()) + .collect() +} + +/// First candidate of `graph` in emit (traversal) order that is in `ties`. +/// Mirrors the reference tie-break: among equal-max-score candidates, take the +/// one encountered first while walking the graph (== HuggingFace merge order). +fn first_in_ties(graph: &GraphV, ties: &FxHashSet>) -> Option> { + crate::graph::first_merge_in(graph, ties) +} + +/// Pick the best merge from `map`, breaking ties by first emit-order appearance +/// while scanning `graphs` in order (the reference's traversal order). +fn pick_best_by_scan<'a>( + map: &FxHashMap, usize>, + graphs: impl Iterator, +) -> Option> { + let (best_key, tie, best) = best_and_tie(map)?; + if !tie { + return Some(best_key.clone()); + } + let ties = tied_at(map, best); + let mut graphs = graphs; + graphs.find_map(|g| first_in_ties(g, &ties)) +} + +/// Tie-break for the connected streaming path: scan documents in order, each +/// assembled on the fly, and take the first tied candidate. +fn pick_best_docs( + map: &FxHashMap, usize>, + doc_words: &[Vec], + cache: &HashMap, +) -> Option> { + let (best_key, tie, best) = best_and_tie(map)?; + if !tie { + return Some(best_key.clone()); + } + let ties = tied_at(map, best); + doc_words + .iter() + .find_map(|words| first_in_ties(&build_doc_graph(words, cache, true), &ties)) +} + +/// Tie-break over unique entries in first-occurrence order, using the inverted +/// index to jump straight to the earliest entry holding a tied candidate. +fn pick_best_entries( + global: &FxHashMap, usize>, + entries: &[WordEntry], + index: &FxHashMap, Vec>, +) -> Option> { + let (best_key, tie, best) = best_and_tie(global)?; + if !tie { + return Some(best_key.clone()); + } + let ties = tied_at(global, best); + let mi = ties + .iter() + .filter_map(|c| index.get(c).and_then(|v| v.iter().min()).copied()) + .min()?; + first_in_ties(&entries[mi as usize].graph, &ties) } fn make_token(nodes: &[GraphV]) -> GraphV { @@ -204,12 +271,12 @@ fn train_single( for _ in range_start..num_merges { let mut counts: FxHashMap, usize> = FxHashMap::default(); graph.emit_merges(&mut |m| *counts.entry(m).or_insert(0) += 1); - let Some((nodes, count)) = pick_best(&counts) else { + let Some(nodes) = pick_best_by_scan(&counts, std::iter::once(&*graph)) else { break; }; let token = make_token(&nodes); if verbose { - println!("Merging {:?} count={}", nodes, count); + println!("Merging {:?} count={}", nodes, counts[&nodes]); } *graph = graph.merge(&token, &nodes); apply_merge_to_cluster_cache(&token, &nodes); @@ -270,17 +337,24 @@ fn build_word_entries( doc_words: &[Vec], cache: &HashMap, ) -> Vec { - let mut freq_map: FxHashMap = FxHashMap::default(); + // First-occurrence order (like the reference's dict.fromkeys) so the + // tie-break scan over entries matches traversal order. + let mut order: Vec<&str> = Vec::new(); + let mut freq_map: FxHashMap<&str, usize> = FxHashMap::default(); for words in doc_words { for w in words { - *freq_map.entry(w.clone()).or_insert(0) += 1; + match freq_map.get_mut(w.as_str()) { + Some(f) => *f += 1, + None => { + freq_map.insert(w.as_str(), 1); + order.push(w.as_str()); + } + } } } - freq_map + order .into_iter() - .filter_map(|(w, freq)| { - cache.get(w.as_str()).map(|g| WordEntry::new(g.clone(), freq)) - }) + .filter_map(|w| cache.get(w).map(|g| WordEntry::new(g.clone(), freq_map[w]))) .collect() } @@ -358,12 +432,12 @@ fn train_entries_delta( let mut index = build_candidate_index(entries); for _ in range_start..num_merges { - let Some((nodes, count)) = pick_best(&global) else { + let Some(nodes) = pick_best_entries(&global, entries, &index) else { break; }; let token = make_token(&nodes); if verbose { - println!("Merging {:?} count={}", nodes, count); + println!("Merging {:?} count={}", nodes, global[&nodes]); } let affected: Vec = index.get(&nodes).cloned().unwrap_or_default(); @@ -463,12 +537,12 @@ fn train_streaming_connected( } } - let Some((nodes, count)) = pick_best(&global) else { + let Some(nodes) = pick_best_docs(&global, doc_words, cache) else { break; }; let token = make_token(&nodes); if verbose { - println!("Merging {:?} count={}", nodes, count); + println!("Merging {:?} count={}", nodes, global[&nodes]); } for graph in cache.values_mut() { @@ -538,7 +612,7 @@ fn train_streaming_with_counts( for i in range_start..num_merges { let global = build_global_counts(&entries); - let Some((nodes, _)) = pick_best(&global) else { break }; + let Some(nodes) = pick_best_by_scan(&global, entries.iter().map(|e| &e.graph)) else { break }; let token = make_token(&nodes); entries.par_iter_mut().for_each(|entry| { @@ -606,7 +680,7 @@ fn train_streaming_with_counts( } } - let Some((nodes, _)) = pick_best(&global) else { break }; + let Some(nodes) = pick_best_docs(&global, doc_words, &cache) else { break }; let token = make_token(&nodes); for graph in cache.values_mut() { @@ -757,7 +831,10 @@ impl Trainer { for i in range_start..num_merges { count_merges_into(subs, &active, &mut counts); - let Some((nodes, _)) = pick_best(&counts) else { + let Some(nodes) = pick_best_by_scan( + &counts, + subs.iter().zip(active.iter()).filter(|(_, &a)| a).map(|(g, _)| g), + ) else { break; }; let token = make_token(&nodes); @@ -777,7 +854,7 @@ impl Trainer { for i in range_start..num_merges { counts.clear(); graph.emit_merges(&mut |m| *counts.entry(m).or_insert(0) += 1); - let Some((nodes, _)) = pick_best(&counts) else { + let Some(nodes) = pick_best_by_scan(&counts, std::iter::once(&*graph)) else { break; }; let token = make_token(&nodes); diff --git a/tests/tokenizers/test_tie_break.py b/tests/tokenizers/test_tie_break.py new file mode 100644 index 0000000..4d0acb7 --- /dev/null +++ b/tests/tokenizers/test_tie_break.py @@ -0,0 +1,43 @@ +from collections import Counter + +from complex_tokenization.tokenizer import BNETokenizer, BPETokenizer + +# A deliberately tie-heavy corpus: at the first BPE merge, seven candidate pairs +# share the top score, so the winner is decided purely by the tie-break rule +# (first candidate in traversal/emit order, matching the reference trainer and +# HuggingFace). Expected lists are generated from the reference implementation +# and pin the contract for both implementations (fast aliases this module). +CORPUS = ["ab cd ab cd ef gh ef gh"] * 3 + + +class TestTieBreak: + def test_corpus_actually_has_score_ties(self): + tok = BPETokenizer() + trainer = tok.make_trainer(CORPUS) + counts = Counter(trainer.graph.get_merges()) + best = max((len(m) - 1) * c for m, c in counts.items()) + tied = [m for m, c in counts.items() if (len(m) - 1) * c == best] + assert len(tied) > 1, f"corpus must produce a score tie, got {tied}" + + def test_bpe_tie_break(self): + merges = BPETokenizer().train(CORPUS, num_merges=10) + assert merges == [ + ("a", "b"), + (" ", "c"), + (" c", "d"), + (" ", "e"), + (" e", "f"), + (" ", "g"), + (" g", "h"), + (" ", "ab"), + ] + + def test_bne_tie_break(self): + merges = BNETokenizer(n=4).train(CORPUS, num_merges=10) + assert merges == [ + (" ", "c", "d"), + (" ", "e", "f"), + (" ", "g", "h"), + ("a", "b"), + (" ", "ab"), + ]