From 09d44638d014d167c10dce7b5807042cded12ba7 Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Mon, 4 May 2026 19:51:35 +0200 Subject: [PATCH] feat(io-core): SpecialTokenSplitter decorator + SentencePiece HF JSON gaps MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three connected fixes for chat-template SentencePiece models like Gemma 4: 1. New SpecialTokenSplitter decorator that wraps any Tokenizer and adds longest-match atomic special-token splitting. encode() walks left-to-right matching the registered specials and delegates the gaps to base.encode(); decode() groups contiguous non-special ids and calls base.decode() per run, so byte-level UTF-8 spanning inside SentencePiece's decode is preserved (special-token boundaries always sit on UTF-8 boundaries because specials are literal strings). 2. SentencePieceTokenizer.fromTokenizerJson now reads two pieces of metadata that HF tokenizer.json files carry but the previous implementation ignored: - addSpacePrefix is detected from the normalizer block (a Prepend step with prepend == "▁" means add_dummy_prefix=true). Models like Gemma 4 ship a normalizer that omits this step and the model is trained accordingly; the previous default of true produced leading-space tokens (▁Hi instead of Hi) that diverged from HuggingFace reference IDs. - bosTokenId / eosTokenId resolved from added_tokens by matching common chat-template content names (, <|begin_of_text|>, for BOS; symmetric for EOS). HF tokenizer.json files don't carry dedicated bos/eos fields, so this name-based heuristic mirrors how `transformers` itself looks them up at load time. 3. TokenizerFactory now wraps the SentencePiece result in SpecialTokenSplitter when: - GGUF: tokenizer.ggml.token_type contains CONTROL (3) or USER_DEFINED (4) entries (treated symmetrically — both are atomic chat-template markers; is typically CONTROL, app-defined markers like <|tool_call> are often USER_DEFINED). - tokenizer.json: added_tokens carries entries with "special": true (or with the field absent — defaults to true, matching HuggingFace's own default). Bare tokenizers without specials are returned unwrapped, so the vanilla LLaMA / TinyLlama path sees no behavior change. The decorator is intentionally generic and can later subsume the inline special-token logic in QwenByteLevelBpeTokenizer and TekkenTokenizer (a follow-up — leaving them alone here keeps the blast radius minimal). Adds SpecialTokenSplitterTest covering empty specials, single special between segments, longest-match priority, consecutive specials, and bos/eos delegation; extends TokenizerFactoryDispatchTest with two new cases proving SP wraps when specials are present (GGUF token_type == 3 path and tokenizer.json added_tokens path). Closes a real downstream bug surfaced in SKaiNET-transformers: Gemma 4 SafeTensors loading via tokenizer.json fragmented chat- template markers (, <|turn>, ) into per-character byte fallbacks, breaking generation. With this change the GGUF path picks up the same atomic-marker behavior automatically, and the HF JSON path matches HuggingFace transformers byte-for-byte. Co-Authored-By: Claude Opus 4.7 (1M context) --- .../io/tokenizer/SentencePieceTokenizer.kt | 66 ++++++++++ .../io/tokenizer/SpecialTokenSplitter.kt | 119 ++++++++++++++++++ .../sk/ainet/io/tokenizer/TokenizerFactory.kt | 76 ++++++++++- .../io/tokenizer/SpecialTokenSplitterTest.kt | 100 +++++++++++++++ .../tokenizer/TokenizerFactoryDispatchTest.kt | 38 ++++++ 5 files changed, 395 insertions(+), 4 deletions(-) create mode 100644 skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitter.kt create mode 100644 skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitterTest.kt diff --git a/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SentencePieceTokenizer.kt b/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SentencePieceTokenizer.kt index d6cbe2e9f..16d711cb6 100644 --- a/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SentencePieceTokenizer.kt +++ b/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SentencePieceTokenizer.kt @@ -1,6 +1,8 @@ package sk.ainet.io.tokenizer import kotlinx.serialization.json.JsonObject +import kotlinx.serialization.json.boolean +import kotlinx.serialization.json.booleanOrNull import kotlinx.serialization.json.doubleOrNull import kotlinx.serialization.json.float import kotlinx.serialization.json.int @@ -319,6 +321,21 @@ public class SentencePieceTokenizer( * * HF Unigram stores the vocab as a JSON array of `[token, score]` * pairs, indexed by id. The unknown token id is at `model.unk_id`. + * + * `addSpacePrefix` is detected from the `normalizer` block — a + * `Prepend` step with `prepend == "▁"` (or a normalizer Sequence + * containing one) means the model expects the SP whitespace + * dummy prefix. Gemma 4 omits it (HF normalizer is `Replace`-only) + * and the model is trained accordingly; without honoring this + * flag the encoder produces leading-space tokens (`▁Hi` instead + * of `Hi`) that diverge from HuggingFace reference IDs. + * + * `bosTokenId` / `eosTokenId` are resolved from `added_tokens` by + * matching common chat-template names (`` / `<|begin_of_text|>` / + * `` for BOS; `` / `<|end_of_text|>` / `` for EOS). + * HF tokenizer.json files do not carry a dedicated bos/eos field, + * so this name-based heuristic mirrors how `transformers` itself + * looks them up at load time. */ public fun fromTokenizerJson(root: JsonObject): SentencePieceTokenizer { val model = root["model"]?.jsonObject @@ -336,13 +353,62 @@ public class SentencePieceTokenizer( } val unknownId = model["unk_id"]?.jsonPrimitive?.int + val addSpacePrefix = detectAddSpacePrefix(root) + val (bosId, eosId) = extractBosEosFromAddedTokens(root) + return SentencePieceTokenizer( tokens = tokens, scores = scores, unknownTokenId = unknownId, + bosTokenId = bosId, + eosTokenId = eosId, + addSpacePrefix = addSpacePrefix, ) } + /** + * Inspect `tokenizer.json#normalizer` for an SP-style `Prepend` + * step with `prepend == "▁"`. Returns true when found, false when + * a normalizer block exists but no such step is present, true as + * a default when there's no normalizer (matches the SPM library + * default of `add_dummy_prefix=true`). + */ + private fun detectAddSpacePrefix(root: JsonObject): Boolean { + val normalizer = root["normalizer"] as? JsonObject ?: return true + val seq = normalizer["normalizers"]?.jsonArray + val candidates = seq ?: listOf(normalizer) + for (n in candidates) { + val obj = n as? JsonObject ?: continue + val type = obj["type"]?.jsonPrimitive?.content ?: continue + if (type == "Prepend") { + val prepend = obj["prepend"]?.jsonPrimitive?.content + if (prepend == "▁") return true + } + } + return false + } + + /** + * Resolve BOS/EOS ids from `added_tokens` by matching standard + * chat-template content names. First match wins; returns null + * for either when no entry matches. + */ + private fun extractBosEosFromAddedTokens(root: JsonObject): Pair { + val added = root["added_tokens"]?.jsonArray ?: return null to null + var bosId: Int? = null + var eosId: Int? = null + for (entry in added) { + val obj = entry as? JsonObject ?: continue + val content = obj["content"]?.jsonPrimitive?.content ?: continue + val id = obj["id"]?.jsonPrimitive?.int ?: continue + when (content) { + "", "<|begin_of_text|>", "" -> if (bosId == null) bosId = id + "", "<|end_of_text|>", "" -> if (eosId == null) eosId = id + } + } + return bosId to eosId + } + private const val TOKEN_TYPE_UNKNOWN = 2 } } diff --git a/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitter.kt b/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitter.kt new file mode 100644 index 000000000..d1cd5d7d4 --- /dev/null +++ b/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitter.kt @@ -0,0 +1,119 @@ +package sk.ainet.io.tokenizer + +/** + * Decorator that adds atomic special-token splitting on top of any base + * [Tokenizer] that does not already implement it. + * + * Why this exists + * --------------- + * Chat-template models embed control markers — ``, `<|im_start|>`, + * `<|turn>`, ``, etc. — that the model is trained to see as a + * single id rather than as the byte/SP-encoded fragments their literal + * string would otherwise produce. [QwenByteLevelBpeTokenizer] and + * [TekkenTokenizer] handle this internally because their input formats + * (GGUF, tekken.json) carry an explicit special-token registry. + * [SentencePieceTokenizer], by design, does not — it implements pure + * llama.cpp-style SPM with byte fallback and nothing else. + * + * That left a gap for HuggingFace-flavored SentencePiece models like + * Gemma 4 whose `tokenizer.json#added_tokens` registry holds the chat- + * template specials. This decorator closes the gap: callers (typically + * [TokenizerFactory]) construct the bare base tokenizer, extract the + * specials map from the source format, and wrap. + * + * The decorator is intentionally generic — a future refactor can lift + * the inline special-token logic out of [QwenByteLevelBpeTokenizer] and + * [TekkenTokenizer] and use this decorator there too. For this change + * we only apply it to SentencePiece, which is the smallest blast-radius + * fix that closes the actual user-visible bug. + * + * Algorithm + * --------- + * - **encode(text)**: walk left-to-right; at each position try the + * longest registered special-token string. On a match, emit its id + * and advance past it. Otherwise extend the current non-special + * segment until the next special boundary (or end-of-text), then + * `base.encode(segment)`. + * - **decode(ids)**: scan ids; collect contiguous non-special ids into + * a buffer, flushing via `base.decode(buffer)` when we hit a special + * id, then emit the special's string form. The byte-level UTF-8 + * spanning that some bases (notably SentencePiece) do inside their + * `decode(IntArray)` is preserved within each non-special run, + * because special-token boundaries always sit on UTF-8 boundaries + * (specials are literal strings). + */ +public class SpecialTokenSplitter( + private val base: Tokenizer, + private val specialTokens: Map, + override val bosTokenId: Int? = base.bosTokenId, + override val eosTokenId: Int? = base.eosTokenId, +) : Tokenizer { + + private val specialIdToString: Map = + specialTokens.entries.associate { (k, v) -> v to k } + + /** Longest-first ordering so e.g. `<|im_start|>` wins over `<|im`. */ + private val specialsByLengthDesc: List = + specialTokens.keys.sortedByDescending { it.length } + + override val vocabSize: Int = base.vocabSize + + override fun encode(text: String): IntArray { + if (specialTokens.isEmpty()) return base.encode(text) + val out = ArrayList(text.length) + var i = 0 + while (i < text.length) { + val matched = matchSpecialAt(text, i) + if (matched != null) { + out.add(specialTokens.getValue(matched)) + i += matched.length + continue + } + val nextSpecial = nextSpecialStart(text, i) + val segment = text.substring(i, nextSpecial) + for (id in base.encode(segment)) out.add(id) + i = nextSpecial + } + return IntArray(out.size) { out[it] } + } + + override fun decode(ids: IntArray): String { + if (ids.isEmpty()) return "" + if (specialTokens.isEmpty()) return base.decode(ids) + val sb = StringBuilder() + val buffer = ArrayList() + for (id in ids) { + val special = specialIdToString[id] + if (special != null) { + if (buffer.isNotEmpty()) { + sb.append(base.decode(buffer.toIntArray())) + buffer.clear() + } + sb.append(special) + } else { + buffer.add(id) + } + } + if (buffer.isNotEmpty()) { + sb.append(base.decode(buffer.toIntArray())) + } + return sb.toString() + } + + private fun matchSpecialAt(text: String, from: Int): String? { + for (tok in specialsByLengthDesc) { + if (tok.isNotEmpty() && text.regionMatches(from, tok, 0, tok.length)) return tok + } + return null + } + + private fun nextSpecialStart(text: String, from: Int): Int { + var earliest = text.length + for (tok in specialTokens.keys) { + if (tok.isEmpty()) continue + val idx = text.indexOf(tok, startIndex = from + 1) + if (idx in 0 until earliest) earliest = idx + } + return earliest + } +} diff --git a/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/TokenizerFactory.kt b/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/TokenizerFactory.kt index 0874696ec..dddb01af7 100644 --- a/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/TokenizerFactory.kt +++ b/skainet-io/skainet-io-core/src/commonMain/kotlin/sk/ainet/io/tokenizer/TokenizerFactory.kt @@ -1,6 +1,10 @@ package sk.ainet.io.tokenizer import kotlinx.serialization.json.Json +import kotlinx.serialization.json.JsonObject +import kotlinx.serialization.json.boolean +import kotlinx.serialization.json.int +import kotlinx.serialization.json.jsonArray import kotlinx.serialization.json.jsonObject import kotlinx.serialization.json.jsonPrimitive import kotlin.jvm.JvmStatic @@ -42,7 +46,10 @@ public object TokenizerFactory { ) return when (model) { "gpt2", "bpe" -> QwenByteLevelBpeTokenizer.fromGgufFields(fields) - "llama", "sentencepiece" -> SentencePieceTokenizer.fromGgufFields(fields) + "llama", "sentencepiece" -> wrapSentencePieceWithSpecialsFromGguf( + base = SentencePieceTokenizer.fromGgufFields(fields), + fields = fields, + ) "bert", "wordpiece" -> throw UnsupportedTokenizerException( "WordPiece/BERT tokenizer not yet implemented" ) @@ -56,8 +63,9 @@ public object TokenizerFactory { * Build a tokenizer from a HuggingFace `tokenizer.json` string. * * Dispatches on `model.type`: `"BPE"` + byte-level pretokenizer routes - * to [QwenByteLevelBpeTokenizer]; `"Unigram"` (SentencePiece) and - * `"WordPiece"` currently throw. + * to [QwenByteLevelBpeTokenizer]; `"Unigram"` (SentencePiece) gets + * wrapped in [SpecialTokenSplitter] when its `added_tokens` registry + * is non-empty; `"WordPiece"` currently throws. */ @JvmStatic public fun fromTokenizerJson(json: String): Tokenizer { @@ -66,7 +74,10 @@ public object TokenizerFactory { ?: throw UnsupportedTokenizerException("tokenizer.json has no model.type") return when (modelType) { "BPE" -> QwenByteLevelBpeTokenizer.fromTokenizerJson(root) - "Unigram" -> SentencePieceTokenizer.fromTokenizerJson(root) + "Unigram" -> wrapSentencePieceWithSpecialsFromJson( + base = SentencePieceTokenizer.fromTokenizerJson(root), + root = root, + ) "WordPiece" -> throw UnsupportedTokenizerException( "WordPiece tokenizer.json not yet implemented" ) @@ -76,6 +87,63 @@ public object TokenizerFactory { } } + /** + * Apply the [SpecialTokenSplitter] decorator to a SentencePiece base + * if the GGUF metadata carries any CONTROL (3) or USER_DEFINED (4) + * token-type entries. Both are atomic chat-template markers that the + * model expects to see as single ids — `` is typically CONTROL, + * `<|tool_call>` and similar app-specific markers are USER_DEFINED. + * The bare base is returned when no specials are present (vanilla + * LLaMA-style models) so consumers see no behavior change. + */ + @Suppress("UNCHECKED_CAST") + private fun wrapSentencePieceWithSpecialsFromGguf( + base: SentencePieceTokenizer, + fields: Map, + ): Tokenizer { + val tokens = (fields["tokenizer.ggml.tokens"] as? List<*>) + ?.filterIsInstance().orEmpty() + val tokenTypes = (fields["tokenizer.ggml.token_type"] as? List<*>) + ?.mapNotNull { (it as? Number)?.toInt() }.orEmpty() + if (tokens.isEmpty() || tokenTypes.isEmpty()) return base + + val specials = HashMap() + val limit = minOf(tokens.size, tokenTypes.size) + for (i in 0 until limit) { + val type = tokenTypes[i] + if (type == TOKEN_TYPE_CONTROL || type == TOKEN_TYPE_USER_DEFINED) { + val tok = tokens[i] + if (tok.isNotEmpty()) specials[tok] = i + } + } + return if (specials.isEmpty()) base else SpecialTokenSplitter(base, specials) + } + + /** + * Apply the [SpecialTokenSplitter] decorator to a SentencePiece base + * built from `tokenizer.json` if its `added_tokens` array carries any + * `"special": true` (or unset, defaulting to true) entries. Returns + * the bare base when the registry is empty. + */ + private fun wrapSentencePieceWithSpecialsFromJson( + base: SentencePieceTokenizer, + root: JsonObject, + ): Tokenizer { + val added = root["added_tokens"]?.jsonArray ?: return base + val specials = HashMap(added.size) + for (entry in added) { + val obj = entry as? JsonObject ?: continue + val content = obj["content"]?.jsonPrimitive?.content ?: continue + val id = obj["id"]?.jsonPrimitive?.int ?: continue + val isSpecial = obj["special"]?.jsonPrimitive?.boolean ?: true + if (isSpecial) specials[content] = id + } + return if (specials.isEmpty()) base else SpecialTokenSplitter(base, specials) + } + + private const val TOKEN_TYPE_CONTROL = 3 + private const val TOKEN_TYPE_USER_DEFINED = 4 + internal val JSON: Json = Json { ignoreUnknownKeys = true; isLenient = true } } diff --git a/skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitterTest.kt b/skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitterTest.kt new file mode 100644 index 000000000..992a7a4ac --- /dev/null +++ b/skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/SpecialTokenSplitterTest.kt @@ -0,0 +1,100 @@ +package sk.ainet.io.tokenizer + +import kotlin.test.Test +import kotlin.test.assertContentEquals +import kotlin.test.assertEquals + +class SpecialTokenSplitterTest { + + /** + * Toy base tokenizer that turns each character into one id (matching its + * code point) and concatenates ids back to chars on decode. Lets us test + * the splitter's segmentation logic without dragging in a real BPE/SP + * vocab. + */ + private object CharBase : Tokenizer { + override val vocabSize: Int = 0x10000 + override val bosTokenId: Int? = null + override val eosTokenId: Int? = null + override fun encode(text: String): IntArray = IntArray(text.length) { text[it].code } + override fun decode(ids: IntArray): String = + buildString(ids.size) { for (id in ids) append(id.toChar()) } + } + + @Test + fun `empty special tokens map is a pass-through`() { + val splitter = SpecialTokenSplitter(CharBase, emptyMap()) + val ids = splitter.encode("hello") + assertEquals(5, ids.size) + assertEquals("hello", splitter.decode(ids)) + } + + @Test + fun `single special token between text segments encodes and decodes atomically`() { + val specials = mapOf("<|end|>" to 99999) + val splitter = SpecialTokenSplitter(CharBase, specials) + val ids = splitter.encode("hi<|end|>bye") + // 'h','i' (2) + special (1) + 'b','y','e' (3) = 6 ids + assertEquals(6, ids.size) + assertEquals(99999, ids[2]) + assertEquals("hi<|end|>bye", splitter.decode(ids)) + } + + @Test + fun `longest match wins over prefix overlap`() { + // Both "<|im" and "<|im_start|>" registered; longer should win. + val specials = mapOf( + "<|im" to 7, + "<|im_start|>" to 8, + ) + val splitter = SpecialTokenSplitter(CharBase, specials) + val ids = splitter.encode("a<|im_start|>b") + // 'a' (1) + special-8 (1) + 'b' (1) = 3 ids + assertContentEquals(intArrayOf('a'.code, 8, 'b'.code), ids) + assertEquals("a<|im_start|>b", splitter.decode(ids)) + } + + @Test + fun `consecutive specials emit no spurious base segments`() { + val specials = mapOf("" to 1, "" to 2) + val splitter = SpecialTokenSplitter(CharBase, specials) + val ids = splitter.encode("") + assertContentEquals(intArrayOf(1, 2), ids) + assertEquals("", splitter.decode(ids)) + } + + @Test + fun `text with no specials at all returns base encode unchanged`() { + val specials = mapOf("" to 42) + val splitter = SpecialTokenSplitter(CharBase, specials) + val ids = splitter.encode("plain") + assertEquals(5, ids.size) + assertEquals("plain", splitter.decode(ids)) + } + + @Test + fun `bos and eos default to base when not overridden`() { + val base = object : Tokenizer { + override val vocabSize: Int = 10 + override val bosTokenId: Int? = 1 + override val eosTokenId: Int? = 2 + override fun encode(text: String): IntArray = IntArray(0) + override fun decode(ids: IntArray): String = "" + } + val splitter = SpecialTokenSplitter(base, emptyMap()) + assertEquals(1, splitter.bosTokenId) + assertEquals(2, splitter.eosTokenId) + } + + @Test + fun `bos and eos overrides take precedence over base`() { + val splitter = SpecialTokenSplitter( + base = CharBase, + specialTokens = emptyMap(), + bosTokenId = 100, + eosTokenId = 200, + ) + assertEquals(100, splitter.bosTokenId) + assertEquals(200, splitter.eosTokenId) + } +} diff --git a/skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/TokenizerFactoryDispatchTest.kt b/skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/TokenizerFactoryDispatchTest.kt index d4e6dfe2e..1789f4d30 100644 --- a/skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/TokenizerFactoryDispatchTest.kt +++ b/skainet-io/skainet-io-core/src/commonTest/kotlin/sk/ainet/io/tokenizer/TokenizerFactoryDispatchTest.kt @@ -120,4 +120,42 @@ class TokenizerFactoryDispatchTest { TokenizerFactory.fromTokenizerJson(json) } } + + @Test + fun `gguf llama with control tokens wraps SentencePiece in splitter`() { + // Vocab: (0,UNK=2) (1,CONTROL=3) (2,CONTROL=3) ▁(3) a(4) + val fields = mapOf( + "tokenizer.ggml.model" to "llama", + "tokenizer.ggml.tokens" to listOf("", "", "", "▁", "a"), + "tokenizer.ggml.scores" to listOf(0.0f, 0.0f, 0.0f, -1.0f, -1.0f), + "tokenizer.ggml.token_type" to listOf(2, 3, 3, 1, 1), + "tokenizer.ggml.bos_token_id" to 1, + "tokenizer.ggml.eos_token_id" to 2, + ) + val tok = TokenizerFactory.fromGguf(fields) + assertTrue(tok is SpecialTokenSplitter, "expected SpecialTokenSplitter wrapping SentencePiece, got ${tok::class.simpleName}") + // in the middle of text encodes as the atomic id 1, not as + // SentencePiece byte-fallback fragments. + val ids = tok.encode("aa") + assertEquals(1, ids.toList().count { it == 1 }, "exactly one id expected, got ${ids.toList()}") + } + + @Test + fun `tokenizer_json Unigram with added_tokens wraps SentencePiece in splitter`() { + val json = """ + { + "added_tokens": [ + {"id": 1, "content": "", "special": true}, + {"id": 2, "content": "", "special": true} + ], + "model": { + "type": "Unigram", + "unk_id": 0, + "vocab": [["", 0.0], ["", 0.0], ["", 0.0], ["▁", -1.0], ["a", -1.0]] + } + } + """.trimIndent() + val tok = TokenizerFactory.fromTokenizerJson(json) + assertTrue(tok is SpecialTokenSplitter, "expected SpecialTokenSplitter wrapping SentencePiece, got ${tok::class.simpleName}") + } }