Files
HomeAgent/internal/memory/qwen/tokenizer.go
JianFeeeee 8e88ae789f feat(memory): 千问文本塔 ONNX 嵌入器(onnxruntime 标签,含 stub)
与 internal/memory/clip 同模式:`//go:build onnxruntime` 编真实实现,无标签时
走 stub,默认构建不链接 onnxruntime、行为不变。

加载契约(目录由 core.memory.multimodal_space.model_dir 指定):
TextTower.onnx + 外部权重分片、tokenizer.json、embed_config.json。
图内已含 last-token 池化,输出即 [batch, dim];L2 归一化在 Go 侧做。

两个刻意的设计选择:

1. **EmbedImageDense 明确报错,不返回零向量**
   导出的是文本塔,视觉塔未导出。返回零向量会让「写入了但检索不到」,
   把跨模态检索失效变成静默故障;明确报错则调用方(mediaref.go)log 后
   跳过写向量,文本路径不受影响。

2. **Fingerprint 只哈希图文件 + 配置 + 外部权重的文件名与大小**
   该目录有 6.5GB 权重分片,启动时全读一遍要几十秒、会阻塞 homeagent 启动。
   换模型必然改变文件集合或大小,足以识别切换;代价是理论上存在
   「大小相同但内容不同」的漏判,对本地单机部署可接受。已写入注释。

模板渲染(renderInstructionInput)放在无构建标签的 tokenizer.go,因此可被
测试覆盖:参考数据里有该模板串的用例,逐 token 对齐验证——模板差一个字符,
池化取到的「最后一个有效 token」位置就变,向量就不同,且不会报错。

验证:6 个测试全绿;go build ./...(stub)与 go build -tags onnxruntime
(真实实现)均通过。
2026-09-11 00:16:45 +08:00

482 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、不做截断
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
}