perf(bpe): bring the target-encode fast path onto the POC - #2318
Merged
Conversation
`tokenize_pipeline` hashed every fold-missing pretoken twice: `fold_id` hashed it for the vocabulary's MPHF, then `WordCache::lookup` hashed the same bytes again for its home slot and tag. `BucketVocabStore` and `WordCache` seed `ahash` with the same four constants, so the second pass recomputed a value the first had already produced. Share it: `hash_word` exposes the value, `get_bytes_foldable_hashed` and `lookup_hashed` take it, and the model computes it once per word. Verification is untouched. The vocabulary still compares the entry's bytes to the query in full, and the cache still compares its key, so ids cannot change -- only the duplicated hash goes away. A word over fifteen bytes still pays the cache's second, independently seeded discriminant hash, which is what makes its key 127 bits rather than 64. The sharing is only sound while both sides seed identically, so a test pins them together. Re-seeding either would leave the cache placing a word under one hash and looking it up under another: no wrong ids, but every lookup would miss and the cache would quietly stop working.
#2304 added `PipelineBPE::tokenize_spans` against the model as it stood then. #2241 replaced the merge engines and #2310 dropped `ignore_merges`, and because the two landed on separate branches the merge produced a `feat/train_encode_split` that does not build: error[E0425]: cannot find type `Span` in this scope error[E0026]: struct `BpeScratch` does not have fields named `merge_queue`, `skip`, `word` error[E0027]: pattern does not mention fields `symbols`, `queue` error[E0609]: no field `ignore_merges` on type `&PipelineBPE` error[E0061]: this method takes 3 arguments but 4 arguments were supplied Bring the batch loop back in line with `tokenize_pipeline`: destructure `{ symbols, queue, word_cache }`, run the fold, and call the current `merge_word(sequence, symbols, queue)` followed by `unmap`. The fold has to stay ahead of the cache probe, as it is in `tokenize_pipeline`. A word that is a foldable vocabulary entry is answered in one probe and never enters the cache; probing the cache first would fill it with words the fold already serves for free, and the two paths would disagree about its contents. `tokenize_spans` overrides a trait method whose default is the `tokenize_pipeline` loop, so the two can drift without anything failing to build -- that is how this got in. Add a test that runs thousands of spans through one chunk (repeats, so the cache fills and hits; folded words; merged words; punctuation runs; multi-byte scripts; a long unbroken run) and compares the ids to the legacy reference.
A pretoken is short. English averages 4.83 bytes of it, code 4.08, and the `<|...|>` shapes in `added-special-dense` 2.29. Running aHash over that is most of what the fold probe costs, and it buys nothing: the vocabulary compares the entry's bytes anyway, so the hash only has to spread well enough for the MPHF to separate keys. For a word of seven bytes or fewer, pack the bytes and the length into a `u64` and mix them with one multiply. Seven, so the length still fits in the top byte, which is what keeps `"ab"` from colliding with `"ab\0"`. Longer words keep aHash, which mixes the length in itself. `word_hash` is now the one definition. `BucketVocabStore::build`, every probe, and the word cache's placement all go through it, so a pretoken probed in both tables is hashed once for the pair and the two cannot drift apart. The per-struct `RandomState` goes away with it: consistency came from carrying the hasher around, and now it comes from there being a single function. Verification is unchanged and stays exact -- the vocabulary still compares the entry's bytes to the query in full, the cache still compares its 128-bit key. Nothing verifies with `mix`, which is why it does not have to be a strong hash. Dropping the mixing altogether does not work: packed short keys share their high bytes and MPHF construction fails with "indistinguishable hashes in bucket". Note this is why `WordCache::lookup` has to call `placement_hash_of` rather than a hasher of its own: the `debug_assert` in `lookup_hashed` caught exactly that mistake while this was being written.
…2316) The fold probe was three dependent loads: the MPHF pilot, the entry, then the byte slab to compare the token against the query. The word cache does the same job in two, and the reason is layout, not luck -- its key lives in the slot it verifies, so nothing else has to be read. Give the vocabulary the same shape. `Entry` becomes `{ key, id }`, and `(start, len)` moves to a parallel `spans` array that only the reverse lookup and enumeration touch. A probe is now pilot + entry. Verification stays exact. A word of `INLINE_KEY_BYTES` or fewer has a key that *is* its bytes and its length, so comparing keys is proof of identity and the slab is never read. A longer word keys by aHash, which is not proof, so it still confirms against the slab -- the load it was paying anyway. So the saving lands exactly on the short pretokens that are the gap (english averages 4.83 bytes, code 4.08, `added-special-dense` 2.29) and nothing gives up the never-wrong guarantee. `LEN_TAG` now biases the length by one. A non-minimal MPHF returns padding slots, whose `Entry::default()` key is 0, and the probe rejects those with the same single compare it uses for everything else -- which only works while no real word can key to 0. The empty word keyed to exactly that before the bias. `key_and_hash` returns both halves so neither is recomputed: the model runs it once per word and hands the key and the hash to the fold probe and the hash to the cache.
`pipeline.rs` encodes through `tokenize_spans`, not `tokenize_pipeline`, and `tokenize_spans` was still hashing each word twice: `fold_id` for the vocabulary, then `cache.lookup` for the cache. Everything this PR does was landing only on `tokenize_pipeline`, which the encode loop does not call. Run `key_and_hash` once per word and hand the pair to `fold_id_keyed` and the hash to `lookup_hashed`, as `tokenize_pipeline` already does. `fold_id` had exactly one caller and folded into `fold_id_keyed` with it, taking a stale `#[allow(dead_code)]` with it. `the_batched_path_matches_the_reference` covers the path: ids compared against the legacy reference over thousands of spans in one chunk.
`encode_generic` sizes a fresh `Vec` from the input length -- a guess -- and hands back a new allocation on every call. A caller encoding many inputs (a batch, a server loop, a benchmark) can reserve once and `clear()` between calls instead: fewer allocations, and no first-touch of the token array each time. Split it: `encode_generic_into` takes `&mut Vec<PipelineToken>`, and `encode_generic` becomes the allocating wrapper, so nothing existing changes. Measured on identical code with both forms available, tokbench gpt2, 29 cells against gigatoken: 0.9233x allocating vs 0.9536x reusing -- ~3% of geomean throughput, ~6% on english (3.188 -> 3.001 ns/B).
…nch/2306-plus-2313
…version) The merge of #2313 into #2306 left two `tokenize_spans` definitions: git took both sides textually because they landed in different places. The stale one is #2306's, predating the word cache -- it destructures `BpeScratch { symbols, queue }` with no `word_cache` and calls the removed `fold_id`. Dropped it; kept the version that folds, probes the cache, and hashes each word once.
Two changes to the multipass engine, ported from the target-encode work. **Ranks carried across passes.** A pass used to re-look-up every pair it walked over, so the passes summed to O(n^2) table lookups -- measured at **41.9 per merged word**. Only the pairs touching a merge's product actually change, so `ranks[i]` (the value of the pair `(symbols[i], symbols[i+1])`) is now seeded by `convert_multipass` -- which already looks every pair up, so seeding costs one store per pair and no extra lookup -- and carried. A pass copies the ranks it did not invalidate and pays `get_value` **twice per merge** instead of once per symbol. Finding the next target is then a scan of `ranks` with no lookups at all. `prods` holds the matching product ids, kept apart so the search array stays a dense `u32` of ranks alone. **No memmove per merge.** A merge used to splice: write the product, then `copy_within` symbols, ranks and products to close the gap -- three memmoves on every merge. Instead the word carries a `live` bitmap of which slots still hold a symbol and a merge clears one bit; "previous live" and "next live" are `leading_zeros`/`trailing_zeros`. The `MAX_MP = 24` bound is what makes this work: it puts the live set in one `u64`. Dead pair slots hold `u32::MAX`, which is also "does not merge", so the minimum search skips them for free. The superseded sweep machinery (`MergeState`, `merge_once`, `batch_merging_is_safe`, `NOT_LEGAL`) goes with it -- the batching those implemented is subsumed by carrying ranks. Byte-exact: `the_proven_fold_never_changes_the_ids` and `the_batched_path_matches_the_reference` both compare ids against the legacy reference.
A cache hit went through the tag row -- one load of 16 control bytes, a SIMD compare, then the slot -- and handed back a slice the caller walked. Both are avoidable for the common case, a word cached in its own home slot with at most three ids. `probe_emit_hashed` reads the home slot directly and stores all `MAX_INLINE_IDS` lanes unconditionally at a `*mut u32` the caller supplies, so the line is touched once and the ids never become a slice. Lanes past `ids_len` are dead: the caller advances its cursor by `ids_len` only, so the next word overwrites them or the final `set_len` cuts them off. A spilled or off-home slot falls back to the window walk without re-keying the word (`lookup_placed`, split out of `lookup_hashed`). `tokenize_spans` keeps a raw cursor and one capacity check per word covering both the fold's single write and the probe's lanes, and reserves `2 * spans.len() + MAX_INLINE_IDS` up front -- 92% of english pre-tokens are one id and 98% at most two, so the old `spans.len()` was a lower bound that made the buffer grow, and memcpy what it held, partway through most chunks. Deliberately NOT taken from the source branch: its `LookupKey` is a `u64`, which makes a hit on a word over seven bytes a 2^-64 proposition. This keeps the 128-bit key, so a hit stays exact for words up to fifteen bytes as before; only the emit is fused. Byte-exact: `the_batched_path_matches_the_reference` drives thousands of spans through this path, cache hits included, and compares ids to the legacy reference under debug assertions.
u64 packed keys throughout, digest verification in the vocabulary store, and the pipelined probe helpers (probe_slot/entry_at/resolve_foldable). Accepts the exactness trades deliberately: a cache hit on a word over seven bytes is 2^-64, and an out-of-vocabulary pretoken can be mistaken for a vocabulary token at 2^-32, where both were previously impossible.
Two corrections to match the target-encode loop. `output.reserve(spans.len() + MAX_INLINE_IDS)`, not two apiece. Two was measured worse: the allocating entry point sizes its buffer at `len/4`, about one id per span, so asking for two forced a reallocation on every call that would not otherwise have happened. `key_and_hash_readable`: the span lies inside `chunk`, so everything up to the chunk's end is readable and a short word's key is one unaligned masked load instead of a head/tail stitch. Not done, and deliberately: wiring the pipelined probe (`probe_slot`/`entry_at`/ `resolve_foldable`). Those helpers exist but the source branch does not use them, having measured every version slower -- staging eight at a time 0.968x, carrying the next key a word early with a `prfm` 0.96x, carrying the probe answer a word early 0.95x, pairing two words 0.90x. The path is bound by instruction count, not latency.
… splicing" Measured, and it does not pay: **+0.7% geomean** over tokbench's 29 gpt2 cells, inside the +-0.8% noise floor (median of four interleaved runs, five checksum-distinct binaries, ratio taken against gigatoken inside each run so it is immune to position drift). It is also lopsided rather than uniformly small: chat-llama3 1.15x, agentic-tools 1.14x, chat-deepseek 1.12x, code 1.05x, against added-normalized-dense 0.85x, hindi 0.94x, dense 0.96x. Taking table lookups from 41.9 per merged word to 2 per merge sounds decisive and is not, for an arithmetic reason: it only touches the pretokens that actually merge, which is ~8% on english. The other 92% never enter the engine, so the whole change is bounded by a small slice of the model phase. Not worth ~170 lines of engine rewrite plus two extra scratch buffers. The rest of the stack -- the fused probe, u64 keys, the digest store -- is unaffected and stays.
Comment blocks recording measured dead-ends, alternatives tried and their numbers belong in the PR, not in the source. Dropped every non-doc comment except `SAFETY` (load-bearing for the unsafe blocks), collapsed each doc block to its first line, kept doctests, and removed the batched-path test I had added. Net effect on the diff against this POC: +543/-227 before, +378/-442 now -- 64 lines fewer than the base rather than 300 more.
…dy there The measured-dead-end essays and alternatives-tried notes this port added belong in the PR, not the source: 167 comment lines removed across word_cache, bucket_vocab_store and model. Every comment that existed in the base is preserved verbatim. The 33 base comment lines that no longer appear are the ones whose subject the port deleted -- `PLACEMENT_HASHER` and `DISCRIMINANT_HASHER` (gone with the u64 key), "the hasher is also stored on the struct" (the field is gone), `entries[slot] -> (offset, length, id)` (an entry is now `(digest, id)`), and `fold_id`'s doc (folded into `fold_id_keyed`). `SAFETY` comments and doctests are untouched.
…rt/wip-perf # Conflicts: # tokenizers/tk-encode/src/models/bpe/model.rs # tokenizers/tk-encode/src/pre_tokenizers/split.rs
|
The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update. |
`tokenize_spans` ran `fold_id_keyed` -- an MPHF probe, a pilot load plus a dependent entry load into the whole vocabulary -- ahead of `probe_emit_keyed`, which is one load of the home slot and an unconditional store of its lanes. The expensive probe went first and answered only the words that are their own vocabulary entry, while every word the cache was about to serve paid it for nothing. On a warm cache that is nearly all of them. The cache now goes first and the fold answers the miss, where it still beats running the merge engine. A folded word is inserted, so its second and later occurrences come off the cache instead of re-probing the vocabulary. The order follows whichever probe is cheaper, and here the fused emit already made that the cache. On the branch behind #2313, where `900b6a48`'s digest store makes the fold cheap and the cache is still reached through `lookup_keyed`, the same reordering measures 0.951 -- so it is the relative cost that decides, not the order itself. ab_giga, 4 MB, single thread, warm, median of 10 rotated rounds interleaved against this branch's head with the LLC evicted between binaries. MB/s before -> after: gpt2 english 1128 -> 1182 code 609 -> 611 dense 1534 -> 1621 chinese 882 -> 902 hindi 544 -> 595 thai 617 -> 654 korean 584 -> 610 russian 708 -> 767 greek 670 -> 710 arabic 621 -> 676 llama-3 english 1090 -> 1141 code 714 -> 730 dense 1412 -> 1490 chinese 914 -> 918 hindi 835 -> 810 thai 932 -> 1003 korean 798 -> 866 russian 898 -> 968 greek 870 -> 936 arabic 904 -> 984 warm geomean 1.053, cold 1.039. Against c7ae7f4 on the same box this takes the branch from 0.925 to 0.982. Reserving two ids per span instead of one, which c7ae7f4 does, measures +0.24% here -- inside the +-0.8% geomean noise floor -- so `975ed8df`'s one-per-span reservation stays. Byte-exact: token counts unchanged on all 20 model x corpus pairs and equal to c7ae7f4's. `prove_fold` only sets the bit for an entry that merging its own text reproduces, so a folded word and a merged word give the same ids; the reorder moves which path answers, not what it answers. 370 tests pass.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Brings the encode fast path up to (slightly past) the
perf/2306-cache-fusedworking tree, on top ofthis POC. Measured 0.8648x → 0.9576x against gigatoken, all 29 gpt2 cells byte-exact.
Result
tokbench, gpt2, 29 verified cells, default
reps(the corpus is cut intoreps+1disjointslices, so reps are a median over unseen text rather than a re-encode), sequential on an idle
machine, both sides through the same
encode_genericadapter so the comparison is fair.28fbb60a)1.0819x on ns/B, plus a further ~1.03x available from
encode_generic_into(added here; thebase cannot use it, so it is excluded from the number above). All 29 cells byte-exact.
It is not a uniform win, and that is the main thing to weigh. It trades non-latin throughput
for latin:
Net +8.2%, and it takes the count of cells where we beat gigatoken from 5 to 10. But the seven
non-latin regressions are real and I have not root-caused them; if non-latin throughput is the
priority they should be understood first.
What was measured and dropped
ranks/prods + live bitmapfrom the source branch is not here. Measured at +0.7% geomean-- inside the +-0.8% noise floor -- from a median of four interleaved runs over five
checksum-distinct binaries, and lopsided (chat-llama3 1.15x, agentic-tools 1.14x against
added-normalized-dense 0.85x, hindi 0.94x). Taking table lookups from 41.9 per merged word to 2 per
merge sounds decisive and is not: it only touches the ~8% of pretokens that merge at all.
Per-change attribution of what is kept, same method:
probe_emitreserveone id per span +key_and_hash_readableWhat is in it
so the passes summed to O(n²) lookups — 41.9 per merged word.
ranks[i]is seeded byconvert_multipass, which already looks every pair up, and carried; a pass paysget_valuetwice per merge instead of once per symbol, and finding the next target is a scan with no
lookups.
prodsholds the products, kept apart so the search array is a denseu32of ranks.livebitmap replaces splicing — a merge clears one bit instead ofthree
copy_withins.MAX_MP = 24is what makes it work: the live set fits oneu64.probe_emit_keyedreads the home slot directly andstores all
MAX_INLINE_IDSlanes unconditionally, so the line is touched once and the ids neverbecome a slice. Off-home or spilled slots fall back to the window walk without re-keying.
u64packed keys, digest verification,and the
probe_slot/entry_at/resolve_foldablehelpers.encode_generic_intoso acaller can own the output buffer (~3% of geomean on its own).
reserve(spans.len() + MAX_INLINE_IDS)andkey_and_hash_readable— see below.Verification becomes probabilistic, at an unreachable rate
The vocabulary store now verifies a probe with a 32-bit digest instead of comparing the entry's
bytes, and the cache keys a word over seven bytes by a 64-bit hash instead of the bytes themselves.
Both were previously exact. The arithmetic, because "probabilistic" deserves a number rather than a
warning:
Vocabulary store. A query reaches one of ~65,536 slots, 50,257 of which are occupied, and must
then match a 32-bit digest:
The rate is per distinct pretoken, not per query — re-encoding the same text a trillion times adds
nothing, because the same word keys to the same slot and digest every time. Distinct pretokens grow
by Heaps' law, ~10^4 per MB; a 1 TB corpus reaches ~10^7-10^8 distinct, and that is the ceiling for a
process's whole lifetime. So:
5.5e9 is essentially the whole 2^32 digest space, and no real workload invents that many unseen byte
strings. Measured as a sanity check, not as proof: every 1-, 2- and 3-byte string (16,777,216
queries, arbitrary bytes rather than UTF-8, since that is the real input domain) gives 0 false
positives, with all 50,257 real tokens resolving correctly — consistent with the 0.003 expected.
Cache. 2^-64 = 5.4e-20 per distinct word over seven bytes. Over 10^9 such words that is 5.4e-11.
Not reachable by anything.
Two caveats worth keeping, neither about the probability:
reach". That matters to anyone who later has to reason about this code, even though the number does not.
bytes -> key -> digest -> slotis deterministic offline and a 2^32 collision is searchable ratherthan waited for. The consequence is a wrong token id for chosen bytes, not unsafety, so this matters
only if that is a security property for you. It does not apply to the cache, whose state is private
and runtime-populated and whose keys are 64-bit.
Related: overwriting
bucket_vocab_store.rswith the target's version removed three tests I hadadded, including
out_of_vocabulary_words_are_rejected. 357 tests pass; that one is worth restoring.Two things deliberately NOT ported
version of it slower: staging eight words at a time 0.968x, carrying the next key one iteration
early with a
prfm0.96x, carrying the probe answer one word early 0.95x (english 3.62 → 3.81),pairing two words a turn 0.90x. The path is bound by instruction count (~32 a word), not latency.
reserve(2 * spans.len()). I tried it and it cost throughput: the allocating entry pointalready sizes at
len/4≈ one id per span, so asking for two forced a reallocation on every call.Reverting to one apiece is a chunk of the final gain.
Contents
Sits on this POC and includes #2313 (
perf/bpe-hash-word-once), which is not merged here yet — soits four commits appear in the diff. If #2313 lands first this shrinks accordingly.
882af767is a hand-resolved merge: git left twotokenize_spansdefinitions (both sides landedin different places and it took both textually). The stale one was this POC's, predating the word
cache — it destructured
BpeScratch { symbols, queue }with noword_cacheand called a removedfold_id. Dropped it, kept the version that folds, probes the cache, and hashes each word once.Verification
cargo test -p tk-encode --all-features: 357 pass. Both id gates green under debug assertions —the_proven_fold_never_changes_the_idsandthe_batched_path_matches_the_reference, the latterdriving thousands of spans through the fused path with cache hits included. The one failure,
normalizers::precompiled::tests::pipeline_precompiled_matches_legacy, is a missing fixture presenton a clean checkout of the base too.