diff --git a/fast/src/trainer.rs b/fast/src/trainer.rs index 2eea101..ae85c7d 100644 --- a/fast/src/trainer.rs +++ b/fast/src/trainer.rs @@ -132,7 +132,7 @@ fn best_and_tie_sharded( fn pick_best_entries( global: &[FxHashMap, usize>], entries: &[WordEntry], - index: &[FxHashMap, Vec>], + index: &[FxHashMap, FxHashSet>], ) -> Option> { let (best_key, tie, best) = best_and_tie_sharded(global)?; if !tie { @@ -252,74 +252,95 @@ fn apply_premerge( token: &GraphV, merge: &[GraphV], ) { - let results: Vec<(usize, Option)> = registry + // Only words that merge land in `changed`; still-registered pointers are + // known unchanged, so no full result map is built or scanned per call. + let changed: FxHashMap = registry .par_iter() - .map(|(&p, g)| (p, g.try_merge(token, merge))) + .filter_map(|(&p, g)| g.try_merge(token, merge).map(|new_g| (p, new_g))) .collect(); - let word_result: FxHashMap> = results.iter().cloned().collect(); - for (p, r) in results { - if let Some(new_g) = r { - registry.remove(&p); - if let Some(np) = crate::graph::arc_ptr(&new_g) { - registry.insert(np, new_g); - } + for (p, new_g) in &changed { + registry.remove(p); + if let Some(np) = crate::graph::arc_ptr(new_g) { + registry.insert(np, new_g.clone()); } } + let registry = &*registry; let m = merge.len(); subs.par_iter_mut().for_each(|doc| { - let GraphV::Seq(nodes_arc) = &*doc else { + let GraphV::Seq(nodes_arc) = doc else { if let Some(new_doc) = doc.try_merge(token, merge) { *doc = new_doc; } return; }; - let nodes: &[GraphV] = nodes_arc; - let n = nodes.len(); + let n = nodes_arc.len(); - let mut out: Vec = Vec::with_capacity(n); + // Cheap scan first: top-level matches and per-child replacements. + // Untouched docs return without cloning anything. let mut top_match = false; let mut i = 0; while i + m <= n { - if nodes[i..i + m] == *merge { - out.push(token.clone()); - i += m; + if nodes_arc[i..i + m] == *merge { top_match = true; - } else { - out.push(nodes[i].clone()); - i += 1; + break; } - } - while i < n { - out.push(nodes[i].clone()); i += 1; } - - if out.len() == 1 { - let single = out.pop().unwrap(); - if top_match { - *doc = single; - } else if let Some(new_doc) = single.try_merge(token, merge) { - *doc = new_doc; - } - return; - } - - let mut child_changed = false; - for c in out.iter_mut() { + let mut replacements: Vec<(usize, GraphV)> = Vec::new(); + for (i, c) in nodes_arc.iter().enumerate() { if matches!(c, GraphV::Node(_)) { continue; } - let merged = match crate::graph::arc_ptr(c).and_then(|p| word_result.get(&p)) { - Some(cached) => cached.clone(), + let merged = match crate::graph::arc_ptr(c) { + Some(p) => match changed.get(&p) { + Some(new_g) => Some(new_g.clone()), + None if registry.contains_key(&p) => None, + None => c.try_merge(token, merge), + }, None => c.try_merge(token, merge), }; if let Some(new_c) = merged { - *c = new_c; - child_changed = true; + replacements.push((i, new_c)); + } + } + if !top_match && replacements.is_empty() { + return; + } + + if !top_match { + let nodes = Arc::make_mut(nodes_arc); + for (i, new_c) in replacements { + nodes[i] = new_c; + } + return; + } + + // Top-level match: rebuild like `seq_try_merge` — matches consume + // the ORIGINAL children (a child that only collapses to a matching + // Node during this same merge is not re-matched), remaining children + // take their replacements, and a full collapse returns the single + // element without descending into it. + let nodes: &[GraphV] = nodes_arc; + let replaced: FxHashMap = replacements.into_iter().collect(); + let mut out: Vec = Vec::with_capacity(n); + let mut i = 0; + while i + m <= n { + if nodes[i..i + m] == *merge { + out.push(token.clone()); + i += m; + } else { + out.push(replaced.get(&i).cloned().unwrap_or_else(|| nodes[i].clone())); + i += 1; } } - if top_match || child_changed { + while i < n { + out.push(replaced.get(&i).cloned().unwrap_or_else(|| nodes[i].clone())); + i += 1; + } + if out.len() == 1 { + *doc = out.pop().unwrap(); + } else { *doc = GraphV::new_seq(out); } }); @@ -505,7 +526,7 @@ impl WordEntry { &mut self, token: &GraphV, merge: &[GraphV], - word_result: Option<&FxHashMap>>, + word_result: WordResult<'_>, ) -> ApplyOutcome { use crate::settings::{MAX_MERGE_SIZE, ONLY_MINIMAL_MERGES}; use std::sync::atomic::Ordering; @@ -532,7 +553,7 @@ impl WordEntry { &mut self, token: &GraphV, merge: &[GraphV], - word_result: Option<&FxHashMap>>, + word_result: WordResult<'_>, ) -> Option { let GraphV::Seq(nodes_arc) = &self.graph else { return None }; let nodes: &[GraphV] = nodes_arc; @@ -554,47 +575,69 @@ impl WordEntry { // them. (A step-wide mutex-sharded memo was measured slower than this // lock-free local one.) let mut local_memo = crate::graph::MergeMemo::default(); - let mut new_nodes: Vec = Vec::with_capacity(n); - let mut changed_positions: Vec = Vec::new(); - // old_ptr -> (old child, new child, occurrence count) + // position -> replacement, applied in place at the end (no full + // children Vec rebuild for the common no-top-match case). + let mut replacements: Vec<(usize, GraphV)> = Vec::new(); + // old_ptr -> (old child, new child, occurrence count) for locally + // merged children (no precomputed step delta) let mut changed_ptrs: FxHashMap = FxHashMap::default(); + // old_ptr -> (precomputed word delta, occurrence count) + type SharedDelta<'a> = (&'a [(Vec, isize)], isize); + let mut changed_shared: FxHashMap> = FxHashMap::default(); // ptr-less children (e.g. Tree) are counted per occurrence, matching seq_recount let mut changed_plain: Vec<(GraphV, GraphV)> = Vec::new(); for (i, child) in nodes.iter().enumerate() { if matches!(child, GraphV::Node(_)) { - new_nodes.push(child.clone()); continue; } - let merged = match crate::graph::arc_ptr(child) - .and_then(|p| word_result.and_then(|wr| wr.get(&p))) - { - Some(cached) => cached.clone(), - None => child.try_merge_m(token, merge, &mut local_memo), + let ptr = crate::graph::arc_ptr(child); + // Precomputed step delta available? (word came from the registry) + let mut shared_delta: Option<&[(Vec, isize)]> = None; + let merged = match (ptr, word_result) { + (Some(p), Some((changed, registry))) => match changed.get(&p) { + Some((new_g, deltas)) => { + shared_delta = Some(deltas); + Some(new_g.clone()) + } + // Still registered means this word does not contain the + // merge; a miss on both maps is a diverged (unregistered) + // child that must be merged locally. + None if registry.contains_key(&p) => None, + None => child.try_merge_m(token, merge, &mut local_memo), + }, + _ => child.try_merge_m(token, merge, &mut local_memo), }; - match merged { - Some(new_child) => { - changed_positions.push(i); - match crate::graph::arc_ptr(child) { - Some(p) => { - changed_ptrs - .entry(p) - .and_modify(|e| e.2 += 1) - .or_insert_with(|| (child.clone(), new_child.clone(), 1)); - } - None => changed_plain.push((child.clone(), new_child.clone())), + if let Some(new_child) = merged { + match (ptr, shared_delta) { + (Some(p), Some(deltas)) => { + changed_shared + .entry(p) + .and_modify(|e| e.1 += 1) + .or_insert((deltas, 1isize)); + } + (Some(p), None) => { + changed_ptrs + .entry(p) + .and_modify(|e| e.2 += 1) + .or_insert_with(|| (child.clone(), new_child.clone(), 1)); } - new_nodes.push(new_child); + (None, _) => changed_plain.push((child.clone(), new_child.clone())), } - None => new_nodes.push(child.clone()), + replacements.push((i, new_child)); } } - if changed_positions.is_empty() { + if replacements.is_empty() { return Some(ApplyOutcome::Unchanged); } let mut delta: FxHashMap, isize> = FxHashMap::default(); + for (deltas, w) in changed_shared.values() { + for (cand, d) in *deltas { + *delta.entry(cand.clone()).or_insert(0) += d * w; + } + } for (old_c, new_c, w) in changed_ptrs.values() { let w = *w; old_c.emit_merges(&mut |m| *delta.entry(m).or_insert(0) -= w); @@ -608,8 +651,11 @@ impl WordEntry { // Boundary pairs: only pair positions touching a changed child can // differ. seq_recount emits pair (i, i+1) iff both children are leaf // Nodes (minimal pair mode). + let repl_at: FxHashMap = + replacements.iter().map(|(i, g)| (*i, g)).collect(); + let new_at = |x: usize| -> &GraphV { repl_at.get(&x).copied().unwrap_or(&nodes[x]) }; let mut starts: FxHashSet = FxHashSet::default(); - for &c in &changed_positions { + for &(c, _) in &replacements { if c > 0 { starts.insert(c - 1); } @@ -623,9 +669,10 @@ impl WordEntry { .entry(vec![nodes[p].clone(), nodes[p + 1].clone()]) .or_insert(0) -= 1; } - if new_nodes[p].is_node() && new_nodes[p + 1].is_node() { + let (np, nq) = (new_at(p), new_at(p + 1)); + if np.is_node() && nq.is_node() { *delta - .entry(vec![new_nodes[p].clone(), new_nodes[p + 1].clone()]) + .entry(vec![np.clone(), nq.clone()]) .or_insert(0) += 1; } } @@ -658,11 +705,23 @@ impl WordEntry { } } - self.graph = GraphV::new_seq(new_nodes); + // The scan borrows are done; write the replacements in place. Entry + // graphs are uniquely owned after canonicalization, so make_mut does + // not clone. + let GraphV::Seq(nodes_arc) = &mut self.graph else { unreachable!() }; + let nodes_mut = Arc::make_mut(nodes_arc); + for (i, new_child) in replacements { + nodes_mut[i] = new_child; + } Some(ApplyOutcome::Delta(ops)) } } +// A changed word's replacement graph plus its once-per-step candidate delta, +// and the (changed, registry) pair the apply path resolves children against. +type WordChange = (GraphV, Vec<(Vec, isize)>); +type WordResult<'a> = Option<(&'a FxHashMap, &'a FxHashMap)>; + // (shard, is_decr, left/entered-entry, freq-scaled amount, candidate) — one // global/index count update, pre-routed to its candidate shard. type CandOp = (u8, bool, bool, usize, Vec); @@ -749,35 +808,35 @@ fn build_global_counts(entries: &[WordEntry]) -> Vec, usiz // Inverted index: candidate -> entry indices that currently contain it. // Each entry appears at most once per candidate (candidate keys are unique per entry). -fn build_candidate_index(entries: &[WordEntry]) -> Vec, Vec>> { +fn build_candidate_index(entries: &[WordEntry]) -> Vec, FxHashSet>> { let n_threads = rayon::current_num_threads().max(1); let chunk_size = (entries.len() / n_threads).max(1); - let chunk_maps: Vec, Vec>>> = entries + let chunk_maps: Vec, FxHashSet>>> = entries .par_chunks(chunk_size) .enumerate() .map(|(chunk_idx, chunk)| { let base = (chunk_idx * chunk_size) as u32; - let mut local: Vec, Vec>> = + let mut local: Vec, FxHashSet>> = (0..CAND_SHARDS).map(|_| FxHashMap::default()).collect(); for (offset, entry) in chunk.iter().enumerate() { for cand in entry.candidates.keys() { local[cand_shard(cand)] .entry(cand.clone()) .or_default() - .push(base + offset as u32); + .insert(base + offset as u32); } } local }) .collect(); - let mut index: Vec, Vec>> = + let mut index: Vec, FxHashSet>> = (0..CAND_SHARDS).map(|_| FxHashMap::default()).collect(); index.par_iter_mut().enumerate().for_each(|(s, shard)| { for chunk in &chunk_maps { for (cand, idxs) in &chunk[s] { - shard.entry(cand.clone()).or_default().extend_from_slice(idxs); + shard.entry(cand.clone()).or_default().extend(idxs.iter().copied()); } } }); @@ -814,7 +873,7 @@ fn train_entries_delta( let affected: Vec = index[cand_shard(&nodes)] .get(&nodes) - .cloned() + .map(|set| set.iter().copied().collect()) .unwrap_or_default(); // Merge each registered unique word once per step; docs then resolve @@ -822,22 +881,37 @@ fn train_entries_delta( // (skips any local work); a miss falls back to a local try_merge, so a // stale registry (entries diverged through the Rebuilt path) only // costs time, never correctness. - let word_result: Option>> = + // Only words that actually merge land in `changed`; "still in the + // registry" is the known-unchanged signal, so no per-step full map + // rebuild or full refresh scan is paid. + let changed_words: Option> = registry.as_deref().map(|reg| { reg.par_iter() - .map(|(&p, g)| (p, g.try_merge(&token, &nodes))) + .filter_map(|(&p, g)| { + let new_g = g.try_merge(&token, &nodes)?; + // This word's candidate delta, computed once per step + // and shared by every doc containing the word. + let mut d: FxHashMap, isize> = FxHashMap::default(); + g.emit_merges(&mut |m| *d.entry(m).or_insert(0) -= 1); + new_g.emit_merges(&mut |m| *d.entry(m).or_insert(0) += 1); + let deltas: Vec<(Vec, isize)> = + d.into_iter().filter(|&(_, v)| v != 0).collect(); + Some((p, (new_g, deltas))) + }) .collect() }); - if let (Some(reg), Some(wr)) = (registry.as_deref_mut(), word_result.as_ref()) { - for (p, r) in wr { - if let Some(new_g) = r { - reg.remove(p); - if let Some(np) = crate::graph::arc_ptr(new_g) { - reg.insert(np, new_g.clone()); - } + if let (Some(reg), Some(changed)) = (registry.as_deref_mut(), changed_words.as_ref()) { + for (p, (new_g, _)) in changed { + reg.remove(p); + if let Some(np) = crate::graph::arc_ptr(new_g) { + reg.insert(np, new_g.clone()); } } } + let word_result: WordResult<'_> = match (changed_words.as_ref(), registry.as_deref()) { + (Some(c), Some(r)) => Some((c, r)), + _ => None, + }; // Apply the merge to each affected entry, producing either a ready // decr/incr delta (incremental path) or the old candidate map for a @@ -848,7 +922,7 @@ fn train_entries_delta( // `min()`. let apply_one = |entry: &mut WordEntry| -> Option { let outcome = if use_try { - entry.apply_merge_try_outcome(&token, &nodes, word_result.as_ref()) + entry.apply_merge_try_outcome(&token, &nodes, word_result) } else { let old = std::mem::take(&mut entry.candidates); entry.apply_merge(&token, &nodes); @@ -958,14 +1032,12 @@ fn train_entries_delta( } if flag { if let Some(v) = ishard.get_mut(&cand) { - if let Some(pos) = v.iter().position(|&x| x == idx) { - v.swap_remove(pos); - } + v.remove(&idx); } } } else { if flag { - ishard.entry(cand.clone()).or_default().push(idx); + ishard.entry(cand.clone()).or_default().insert(idx); } *gshard.entry(cand).or_insert(0) += amount; }