Files
HomeAgent/internal/memory/qwen/tokenizer.go
JianFeeeee 96c1d7baae fix(memory): 千问文本塔补齐 tokenizer post_processor 与真实 ONNX 回归
验证 Go 端到端路径时发现:HuggingFace 的 tokenizer.json 带 TemplateProcessing
post_processor,规则是 `$A <|endoftext|>`——即每段输入末尾都会追加一个
`<|endoftext|>`(151643)。它正是图内 last-token 池化的锚点:

- 漏掉它:ONNX 仍能运行(不会报错),但池化取到的是模板末尾的 assistant
  起始符,而非 post token,整条嵌入向量与上游不一致;
- 截断语义:HuggingFace 在 truncation=true 时先把正文截到 maxLen-1,
  再保留末尾 post token(实测 600×"记忆"→ [511 正文][151643])。

改动:
1. Tokenizer 新增 encodeModelInput(text, maxLen):Encode 后追加 post token,
   并在超长时先截到 maxLen-1;缺 <|endoftext|> 直接报错(防静默错误)。
2. Embedder.VectorizeDense 改用 encodeModelInput。
3. 测试:
   - TestEncodeModelInputPostProcessor:短文本+post token;长文本按 512 截断
     且尾 token 为 post token(用 600×"记忆"确保真的触发截断分支)。
   - embedder_onnx_test.go(onnxruntime 标签):用 Python onnxruntime 1.28
     生成的冻结参考向量验证完整 Go 路径(模板渲染→BPE→ONNX→L2 normalize),
     逐维 diff ≤ 2e-5;同时断言模型输入恰好 23 token 且末尾是 post token。
     产物不在时跳过(与 tokenizer_test 相同约定)。
   - 修正先前测试误用 400×"记忆":BPE 把"记忆"合并为单 token,400 次只有
     400 token 不触发截断;改用 600 次后确实走到 maxLen-1 分支。

其余(ORT 库路径、指纹纳入 .onnx.data)一并随本提交带上。

验证:go test ./internal/memory/qwen 与 -tags onnxruntime 全绿;
FP32 图与 PyTorch 在短/长/等长批次/真实 padding 批次上余弦均 ≥0.99999994。
2026-09-11 03:09:01 +08:00

504 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// 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
}
// DefaultInstruction 是导出脚本随 embed_config.json 写入的默认指令。
const DefaultInstruction = "Represent the user's input."
// renderInstructionInput 按模型自带的对话模板拼输入(无构建标签,便于测试)。
//
// 必须与 HuggingFace processor 的 apply_chat_template(add_generation_prompt=True)
// 产出完全一致:指令放 system、正文放 user、以 assistant 起始符结尾。差一个
// 特殊 token池化取到的「最后一个有效 token」位置就变了嵌入也就不同——
// 而且不会报错。参考数据集里有该模板串的用例,能逐 token 对齐验证。
func renderInstructionInput(instruction, text string) string {
if instruction == "" {
instruction = DefaultInstruction
}
return "<|im_start|>system\n" + instruction +
"<|im_end|>\n<|im_start|>user\n" + text +
"<|im_end|>\n<|im_start|>assistant\n"
}
// Encode 把文本编码为 token id 序列(识别输入中已有的特殊 token
// 但不执行 tokenizer.json 的 post_processor也不做截断
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
}
// encodeModelInput 执行 TextTower 输入所需的 tokenizer post_processor。
//
// tokenizer.json 的 TemplateProcessing 规则是 `$A <|endoftext|>`HuggingFace
// 在 truncation=true 时先把 A 截到 maxLen-1再保留末尾 post token。漏掉它不会
// 触发 ONNX 错误,却会改变池化位置和整条嵌入向量,因此不能直接用 Encode 的结果。
func (t *Tokenizer) encodeModelInput(text string, maxLen int) ([]int, error) {
postID, ok := t.SpecialID("<|endoftext|>")
if !ok {
return nil, fmt.Errorf("tokenizer.json 缺少 post token <|endoftext|>")
}
if maxLen <= 0 {
return nil, fmt.Errorf("maxLen 必须大于 0")
}
ids := t.Encode(text)
if len(ids) >= maxLen {
ids = ids[:maxLen-1]
}
return append(ids, postID), nil
}
// 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
}