mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-22 18:08:04 +00:00
验证 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。
504 lines
14 KiB
Go
504 lines
14 KiB
Go
// 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
|
||
}
|
||
|
||
// 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|>` 在词表里是一个整体 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
|
||
}
|