mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 01:18:08 +00:00
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:
1376
internal/memory/qwen/testdata/qwen_tokenizer_reference.json
vendored
Normal file
1376
internal/memory/qwen/testdata/qwen_tokenizer_reference.json
vendored
Normal file
File diff suppressed because it is too large
Load Diff
463
internal/memory/qwen/tokenizer.go
Normal file
463
internal/memory/qwen/tokenizer.go
Normal file
@ -0,0 +1,463 @@
|
||||
// Package qwen 实现 Qwen3-VL-Embedding 的字节级 BPE 分词器。
|
||||
//
|
||||
// 为什么不复用 clip 的 tokenizer:CLIP 用的是「小写化 + 空白规整 + 词表 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.IsSpace(Unicode 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|>` 在词表里是一个整体 id(151644),若走 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
|
||||
}
|
||||
161
internal/memory/qwen/tokenizer_test.go
Normal file
161
internal/memory/qwen/tokenizer_test.go
Normal file
@ -0,0 +1,161 @@
|
||||
package qwen
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// modelDir 是本地千问模型目录。不存在则跳过——参考数据已固化在 testdata,
|
||||
// 但分词器本身要从 tokenizer.json 加载词表与 merges(11MB,不入库)。
|
||||
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
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user