feat(memory): 千问字节级 BPE 分词器 + 与上游逐条对齐的回归测试

为「千问嵌入模型导出 ONNX 并内嵌」的 Go 侧准备。CLIP 那套 tokenizer 不能复用:
CLIP 是「小写化 + 空白规整 + 词表 BPE」,千问是 **GPT-2 式字节级 BPE**
(先按字节映射到安全 unicode,再对映射结果做合并),中文与空白输入的切分
完全不同。

实现中撞到两个与上游对齐的坑,都由测试暴露:

1. **`\s+(?!\S)` 的语义依赖正则回溯**,不是「等价于 `\s+`」。
   `\s+` 先贪婪吃完整段空白,发现后面是非空白导致 `(?!\S)` 失败,于是回退
   一个字符,**正好留下末尾一个空白**给前面以 ` ?` / `[^…]?` 开头的分支合并。
   这直接决定切分点:`"   leading"` 会切成 `"  "` + `" leading"`,
   而不是 `"   "` + `"leading"`。RE2 不支持 lookaround,近似改写必然对不上,
   所以改成按分支顺序**显式实现**(含 `\s*[\r\n]+` 的回溯语义)。
   实测:近似改写时 26 条里错 3 条,全部是空白串用例。

2. **Go 的 `\s` 只有 ASCII,且 regexp 不支持二进制属性 `\p{White_Space}`**
   (只支持 script/category,直接写会报 invalid character class range)。
   而上游 Rust regex 的 `\s` 正是 White_Space。改用 `unicode.IsSpace` 作为
   唯一判据,避免全角空格/NBSP/行分隔符的切分点漂移。

另外特殊 token(`<|im_start|>` 等 24 个 AddedToken)必须**整体优先匹配**并
按长度降序,否则会被 BPE 拆成子 token,模型收到的输入就变了——且不会报任何错。

验证:`testdata/qwen_tokenizer_reference.json` 由 HuggingFace 真实 tokenizer
生成(26 条用例,覆盖中/英/中英混排/数字/各类空白形态/标点/emoji/特殊 token/
长文本/空串/单字符边界),Go 实现逐条精确对齐。字节↔unicode 映射另测双射性
(有碰撞会让不同字节编成同一 token,静默产生错误输入)。
This commit is contained in:
JianFeeeee
2026-09-11 00:11:58 +08:00
parent b393b7072c
commit e985151da9
3 changed files with 2000 additions and 0 deletions

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,463 @@
// Package qwen 实现 Qwen3-VL-Embedding 的字节级 BPE 分词器。
//
// 为什么不复用 clip 的 tokenizerCLIP 用的是「小写化 + 空白规整 + 词表 BPE」
// 而千问是 **GPT-2 式字节级 BPE**——先把输入按字节映射到一组可见 unicode
// 再对映射后的字符串做 BPE 合并。两者的预处理不可互换,硬套会在中文和
// 空白较多的输入上产出完全不同的 token。
//
// 与上游HuggingFace tokenizer.json 的 Rust 实现)对齐时的两处坑:
//
// 1. pre_tokenizer 正则里的 `\s+(?!\S)` 是**负向前瞻**Go 的 RE2 不支持
// lookaround。该分支只在「空白一直延伸到串尾」时命中而此时贪婪的
// `\s+` 会匹配完全相同的区间,所以直接删掉该分支即为等价改写。
// 2. Go 的 `\s` 只覆盖 ASCII而 Rust regex 的 `\s` 是 Unicode
// `\p{White_Space}`。不换成 \p{White_Space} 的话全角空格、NBSP、
// 行分隔符等的切分点会与上游不一致。
package qwen
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"unicode"
"unicode/utf8"
)
// 空白判定统一用 unicode.IsSpaceUnicode White_Space 属性)。
//
// 不能用 Go 正则里的 \s——那只覆盖 ASCII也不能写 \p{White_Space}——Go 的
// regexp 只支持 script/category不支持二进制属性会报 invalid character
// class range。上游 Rust regex 的 \s 正是 White_Space所以这里以
// unicode.IsSpace 为准。
// specialToken 是一个 AddedToken以整体形式优先匹配不参与 BPE 拆分。
type specialToken struct {
content string
id int
}
// Tokenizer 是千问的字节级 BPE 分词器。
type Tokenizer struct {
vocab map[string]int
ranks map[string]int
// byteEnc 是 GPT-2 的 byte→unicode 映射:把 0..255 每个字节映到一个
// 「可见且不会与正常文本冲突」的 unicode 码点。因为 BPE 词表基于文本构建,
// 直接放原始字节会与合法 UTF-8 冲突。
byteEnc map[byte]rune
// specials 按 content 长度降序,保证「最长优先」——
// 否则 `<|im_start|>` 可能被 `<|im_` 之类的短 token 先切走。
specials []specialToken
// MaxLen 是嵌入用途的截断上限(与导出脚本的 MAX_LENGTH 一致)。
MaxLen int
}
// tokenizerJSON 只取我们需要的部分。
type tokenizerJSON struct {
Model struct {
Vocab map[string]int `json:"vocab"`
Merges []interface{} `json:"merges"`
} `json:"model"`
AddedTokens []struct {
ID int `json:"id"`
Content string `json:"content"`
Special bool `json:"special"`
} `json:"added_tokens"`
}
// LoadTokenizer 从模型目录加载 tokenizer.json。
func LoadTokenizer(modelDir string) (*Tokenizer, error) {
raw, err := os.ReadFile(filepath.Join(modelDir, "tokenizer.json"))
if err != nil {
return nil, fmt.Errorf("read tokenizer.json: %w", err)
}
var tj tokenizerJSON
if err := json.Unmarshal(raw, &tj); err != nil {
return nil, fmt.Errorf("parse tokenizer.json: %w", err)
}
if len(tj.Model.Vocab) == 0 {
return nil, fmt.Errorf("tokenizer.json 的 model.vocab 为空")
}
ranks := make(map[string]int, len(tj.Model.Merges))
for i, m := range tj.Model.Merges {
// merges 有两种形态:字符串 "a b",或数组 ["a","b"]。
var pair string
switch v := m.(type) {
case string:
pair = v
case []interface{}:
if len(v) == 2 {
a, _ := v[0].(string)
b, _ := v[1].(string)
pair = a + " " + b
}
}
if pair != "" {
if _, seen := ranks[pair]; !seen {
ranks[pair] = i
}
}
}
t := &Tokenizer{
vocab: tj.Model.Vocab,
ranks: ranks,
byteEnc: bytesToUnicode(),
MaxLen: 512,
}
for _, at := range tj.AddedTokens {
if at.Special && at.Content != "" {
t.specials = append(t.specials, specialToken{content: at.Content, id: at.ID})
}
}
// 最长优先,避免短 token 抢走长 token 的前缀。
sort.Slice(t.specials, func(i, j int) bool {
return len(t.specials[i].content) > len(t.specials[j].content)
})
return t, nil
}
// VocabSize 返回词表大小(诊断用)。
func (t *Tokenizer) VocabSize() int { return len(t.vocab) }
// SpecialID 返回特殊 token 的 id不存在时 ok=false。
func (t *Tokenizer) SpecialID(content string) (int, bool) {
for _, s := range t.specials {
if s.content == content {
return s.id, true
}
}
return 0, false
}
// Encode 把文本编码为 token id 序列(不含特殊 token、不做截断
func (t *Tokenizer) Encode(text string) []int {
var ids []int
for _, seg := range t.splitSpecials(text) {
if seg.specialID >= 0 {
ids = append(ids, seg.specialID)
continue
}
ids = append(ids, t.encodeOrdinary(seg.text)...)
}
return ids
}
// seg 是「普通文本」或「已识别的特殊 token」二选一。
type seg struct {
text string
specialID int // -1 表示普通文本
}
// splitSpecials 把输入切成普通片段与特殊 token 片段。
//
// 为什么必须先切:`<|im_start|>` 在词表里是一个整体 id151644若走 BPE
// 会被拆成若干子 token编码结果与上游不一致模型看到的输入也就变了。
func (t *Tokenizer) splitSpecials(text string) []seg {
if len(t.specials) == 0 || text == "" {
return []seg{{text: text, specialID: -1}}
}
var out []seg
for len(text) > 0 {
// 找最靠前的特殊 token 出现位置(同位置取最长)。
bestIdx, bestLen, bestID := -1, 0, -1
for _, s := range t.specials {
i := strings.Index(text, s.content)
if i < 0 {
continue
}
if bestIdx == -1 || i < bestIdx || (i == bestIdx && len(s.content) > bestLen) {
bestIdx, bestLen, bestID = i, len(s.content), s.id
}
}
if bestIdx == -1 {
out = append(out, seg{text: text, specialID: -1})
break
}
if bestIdx > 0 {
out = append(out, seg{text: text[:bestIdx], specialID: -1})
}
out = append(out, seg{specialID: bestID})
text = text[bestIdx+bestLen:]
}
return out
}
// encodeOrdinary 对普通文本做「切分 → 字节映射 → BPE 合并」。
func (t *Tokenizer) encodeOrdinary(text string) []int {
if text == "" {
return nil
}
var ids []int
for _, piece := range t.preTokenize(text) {
// 字节级映射:先把 piece 的 UTF-8 字节逐个映射成 unicode 字符。
var sb strings.Builder
for _, b := range []byte(piece) {
sb.WriteRune(t.byteEnc[b])
}
for _, tok := range t.bpe(sb.String()) {
if id, ok := t.vocab[tok]; ok {
ids = append(ids, id)
}
// 词表里找不到的片段直接丢弃:正常情况不会发生
//(词表覆盖全部 256 个字节级字符),发生即数据有问题。
}
}
return ids
}
// ---- pre_tokenizer ----
//
// 上游是一条正则tokenizer.json 的 pre_tokenizer.pretokenizers[0].pattern
//
// (?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}|
// ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+
//
// **为什么不用一个 Go 正则**:末两个分支里的 `\s+(?!\S)` 是负向前瞻RE2 不
// 支持 lookaround而且它的真实语义依赖**回溯**——`\s+` 先贪婪吃完整段空白,
// 发现后面是非空白导致 `(?!\S)` 失败,于是回退一个字符,正好留下末尾一个
// 空白给前面那些以 ` ?` / `[^…]?` 开头的分支合并。这个“留一个”直接决定
// 切分点(`" leading"` 会切成 `" "` + `" leading"` 而不是 `" "` + `"leading"`
// 近似改写必然对不上,所以按分支顺序显式实现。
func (t *Tokenizer) preTokenize(text string) []string {
var out []string
for len(text) > 0 {
switch {
case matchApostrophe(text) > 0:
n := matchApostrophe(text)
out = append(out, text[:n])
text = text[n:]
case matchWord(text) > 0:
n := matchWord(text)
out = append(out, text[:n])
text = text[n:]
case matchDigit(text) > 0:
n := matchDigit(text)
out = append(out, text[:n])
text = text[n:]
case matchPunct(text) > 0:
n := matchPunct(text)
out = append(out, text[:n])
text = text[n:]
case matchNewline(text) > 0:
n := matchNewline(text)
out = append(out, text[:n])
text = text[n:]
default:
// `\s+(?!\S)|\s+` 合一:空白段。
total, lastStart := wsRun(text)
if total == 0 {
// 兜底:不应到达(分支覆盖全部字符),防御性前进一个 rune。
_, size := utf8.DecodeRuneInString(text)
out = append(out, text[:size])
text = text[size:]
continue
}
n := total
if total < len(text) && lastStart > 0 {
n = lastStart // 后面还有非空白 → 回退掉末尾那一个空白
}
out = append(out, text[:n])
text = text[n:]
}
}
return out
}
func runeAt(s string) (rune, int) { return utf8.DecodeRuneInString(s) }
func isLetter(r rune) bool { return unicode.IsLetter(r) }
func isNumber(r rune) bool { return unicode.IsNumber(r) }
func isWS(r rune) bool { return unicode.IsSpace(r) }
// wsRun 返回开头连续空白段的字节长度,以及最后一个空白 rune 的起始字节位置。
func wsRun(s string) (total, lastStart int) {
lastStart = -1
i := 0
for i < len(s) {
r, size := runeAt(s[i:])
if !isWS(r) {
break
}
lastStart = i
i += size
}
return i, lastStart
}
// matchApostrophe`(?i:'s|'t|'re|'ve|'m|'ll|'d)`
func matchApostrophe(s string) int {
if len(s) == 0 || s[0] != '\'' {
return 0
}
rest := s[1:]
// 各后缀互为前缀关系re/ve/ll/s/t/m/d所以先试长的。
for _, suf := range []string{"re", "ve", "ll", "s", "t", "m", "d"} {
if len(rest) >= len(suf) && strings.EqualFold(rest[:len(suf)], suf) {
return 1 + len(suf)
}
}
return 0
}
// matchWord`[^\r\n\p{L}\p{N}]?\p{L}+`
//
// 注意可选字符**排除** \r \n若吃了可选字符却没有字母跟上整个分支失败
// (与正则的“该分支不匹配”一致,不能把可选字符当已消耗)。
func matchWord(s string) int {
i := 0
if r, size := runeAt(s); r != '\r' && r != '\n' && !isLetter(r) && !isNumber(r) {
i = size
}
r, size := runeAt(s[i:])
if !isLetter(r) {
return 0
}
i += size
for i < len(s) {
r, size := runeAt(s[i:])
if !isLetter(r) {
break
}
i += size
}
return i
}
// matchDigit`\p{N}` —— 只吃**一个**数字。
func matchDigit(s string) int {
if r, size := runeAt(s); isNumber(r) {
return size
}
return 0
}
// matchPunct` ?[^\s\p{L}\p{N}]+[\r\n]*`
//
// 开头是**字面空格**(不是 \s所以只可能吃掉一个 U+0020。
func matchPunct(s string) int {
i := 0
if strings.HasPrefix(s, " ") {
i = 1
}
n := 0
for i+n < len(s) {
r, size := runeAt(s[i+n:])
if isWS(r) || isLetter(r) || isNumber(r) {
break
}
n += size
}
if n == 0 {
return 0
}
i += n
for i < len(s) && (s[i] == '\r' || s[i] == '\n') {
i++
}
return i
}
// matchNewline`\s*[\r\n]+`
//
// 贪婪+回溯的真实语义:`\s*` 先吃完整段空白,`[\r\n]+` 无可匹配而回退,
// 最终停在段内**最后一个** \r 或 \n 之前,再把它之后的连续 \r\n 吃掉。
func matchNewline(s string) int {
total, _ := wsRun(s)
if total == 0 {
return 0
}
last := -1
for j := total - 1; j >= 0; j-- {
if s[j] == '\r' || s[j] == '\n' {
last = j
break
}
}
if last < 0 {
return 0
}
end := last
for end < len(s) && (s[end] == '\r' || s[end] == '\n') {
end++
}
return end
}
// bpe 是标准字节级 BPE反复合并 rank 最小的相邻对,直到无可合并。
func (t *Tokenizer) bpe(word string) []string {
symbols := make([]string, 0, len(word))
for _, r := range word {
symbols = append(symbols, string(r))
}
if len(symbols) < 2 {
return symbols
}
for {
bestRank, bestIdx := -1, -1
for i := 0; i+1 < len(symbols); i++ {
r, ok := t.ranks[symbols[i]+" "+symbols[i+1]]
if !ok {
continue
}
if bestRank == -1 || r < bestRank {
bestRank, bestIdx = r, i
}
}
if bestIdx == -1 {
return symbols
}
merged := symbols[bestIdx] + symbols[bestIdx+1]
symbols = append(symbols[:bestIdx], append([]string{merged}, symbols[bestIdx+2:]...)...)
if len(symbols) < 2 {
return symbols
}
}
}
// bytesToUnicode 是 GPT-2 的字节↔unicode 映射表。
//
// 让每个字节都有一个「安全」的可见码点表示,避免原始控制字节混进 BPE 词表。
// 可打印 ASCII 与拉丁补充区保持原样,其余字节映射到 256 之后的码点。
func bytesToUnicode() map[byte]rune {
bs := make([]int, 0, 256)
for b := int('!'); b <= int('~'); b++ {
bs = append(bs, b)
}
for b := 0xA1; b <= 0xAC; b++ {
bs = append(bs, b)
}
for b := 0xAE; b <= 0xFF; b++ {
bs = append(bs, b)
}
inBS := make(map[int]bool, len(bs))
for _, b := range bs {
inBS[b] = true
}
cs := make([]int, len(bs))
copy(cs, bs)
n := 0
for b := 0; b < 256; b++ {
if inBS[b] {
continue
}
bs = append(bs, b)
cs = append(cs, 256+n)
n++
}
out := make(map[byte]rune, 256)
for i, b := range bs {
out[byte(b)] = rune(cs[i])
}
return out
}

View File

@ -0,0 +1,161 @@
package qwen
import (
"encoding/json"
"os"
"path/filepath"
"testing"
)
// modelDir 是本地千问模型目录。不存在则跳过——参考数据已固化在 testdata
// 但分词器本身要从 tokenizer.json 加载词表与 merges11MB不入库
const modelDir = "/home/newqqagent/models/models/qwen--Qwen3-VL-Embedding-2B/snapshots/master"
type tokenizerRef struct {
VocabSize int `json:"vocab_size"`
Cases []struct {
Text string `json:"text"`
IDs []int `json:"ids"`
Tokens []string `json:"tokens"`
} `json:"cases"`
AddedTokens []struct {
Content string `json:"content"`
ID int `json:"id"`
Special bool `json:"special"`
} `json:"added_tokens"`
}
func loadRef(t *testing.T) *tokenizerRef {
t.Helper()
raw, err := os.ReadFile(filepath.Join("testdata", "qwen_tokenizer_reference.json"))
if err != nil {
t.Fatalf("读取参考数据: %v", err)
}
var ref tokenizerRef
if err := json.Unmarshal(raw, &ref); err != nil {
t.Fatalf("解析参考数据: %v", err)
}
return &ref
}
func loadTokenizer(t *testing.T) *Tokenizer {
t.Helper()
if _, err := os.Stat(filepath.Join(modelDir, "tokenizer.json")); err != nil {
t.Skipf("模型目录不可用,跳过: %v", err)
}
tok, err := LoadTokenizer(modelDir)
if err != nil {
t.Fatalf("LoadTokenizer: %v", err)
}
return tok
}
// 与 HuggingFace 的真实 tokenizer 逐条对齐。
//
// 这是本包唯一的正确性判据:字节级 BPE 的失败模式是「看起来能跑但 token 不同」,
// 而 token 不同会让模型收到完全不同的输入,嵌入自然也就错了——不会报任何错。
// 所以必须拿真实输出对照,不能靠读代码断言。
func TestTokenizerMatchesReference(t *testing.T) {
ref := loadRef(t)
tok := loadTokenizer(t)
if got := tok.VocabSize(); got != ref.VocabSize {
t.Errorf("词表大小 = %d参考 %d", got, ref.VocabSize)
}
failed := 0
for _, c := range ref.Cases {
got := tok.Encode(c.Text)
if !sameIDs(got, c.IDs) {
failed++
t.Errorf("不一致 text=%q\n got %v\n want %v", c.Text, got, c.IDs)
}
}
if failed > 0 {
t.Fatalf("%d/%d 条用例不一致", failed, len(ref.Cases))
}
}
func sameIDs(a, b []int) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
// 特殊 token 必须整体匹配:走 BPE 会被拆成子 token模型看到的输入就变了。
func TestSpecialTokensMatchWhole(t *testing.T) {
ref := loadRef(t)
tok := loadTokenizer(t)
for _, at := range ref.AddedTokens {
if !at.Special {
continue
}
got, ok := tok.SpecialID(at.Content)
if !ok {
t.Errorf("特殊 token %q 未从 tokenizer.json 载入", at.Content)
continue
}
if got != at.ID {
t.Errorf("特殊 token %q id=%d参考 %d", at.Content, got, at.ID)
}
// 单独出现时必须编码成恰好一个 id。
ids := tok.Encode(at.Content)
if len(ids) != 1 || ids[0] != at.ID {
t.Errorf("特殊 token %q 应整体编码为 [%d],实际 %v", at.Content, at.ID, ids)
}
}
}
// 最长优先:`<|im_start|>` 不能被更短的 `<|im_end|>` 之类前缀抢走。
func TestSpecialTokenLongestFirst(t *testing.T) {
tok := loadTokenizer(t)
text := "<|im_start|>user\n你好<|im_end|>"
ids := tok.Encode(text)
startID, _ := tok.SpecialID("<|im_start|>")
endID, _ := tok.SpecialID("<|im_end|>")
if len(ids) == 0 || ids[0] != startID {
t.Fatalf("应以 <|im_start|>(%d) 开头,实际 %v", startID, ids)
}
if last := ids[len(ids)-1]; last != endID {
t.Fatalf("应以 <|im_end|>(%d) 结尾,实际 %v", endID, ids)
}
}
// 空串与单字符边界。
func TestTokenizerEdgeCases(t *testing.T) {
tok := loadTokenizer(t)
if got := tok.Encode(""); len(got) != 0 {
t.Errorf("空串应产出 0 个 token实际 %v", got)
}
for _, s := range []string{"a", "中", "1", " "} {
if got := tok.Encode(s); len(got) == 0 {
t.Errorf("%q 应至少产出 1 个 token", s)
}
}
}
// byteEnc 必须是双射256 个字节映射到 256 个互不相同的码点。
// 有碰撞就会让不同字节编成同一个 token静默产生错误输入。
func TestBytesToUnicodeBijective(t *testing.T) {
m := bytesToUnicode()
if len(m) != 256 {
t.Fatalf("映射应覆盖 256 个字节,实际 %d", len(m))
}
seen := map[rune]byte{}
for b, r := range m {
if prev, dup := seen[r]; dup {
t.Fatalf("码点冲突:字节 %d 与 %d 都映射到 %q", prev, b, r)
}
seen[r] = b
}
}