mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +00:00
问题:cmd/homed 里 `case "onnx": qwen.New(modelDir)` 把模型适配写进了核心, `type=onnx` 名义上是格式、实际写死了一个模型家族;2117 行 Qwen 专属代码 (BPE、chat template、M-RoPE、Vision_gN 命名)住在内核树里,还带着一对 `//go:build onnxruntime` 的 stub。加任何新模型都要改内核。 现在核心只认一个模型无关的公共契约(pkg/embedding): - 输入是不透明的 Data+MIME,解码/预处理/时序分组全归 provider - 能力是数据(Info.Modalities),不是接口方法——新增模态无需改核心接口 - 不支持的模态返回 embedding.ErrUnsupportedModality(可 errors.Is 识别) - 按名字注册,重复注册 panic;Options 是 provider 私有命名空间,核心不解释 改动: - 新增 pkg/embedding:Modality/Purpose/Input/Info/Provider/Config + 注册表 (Open 校验 Info,ValidateVector 在入库前拦下维度错与非有限值) - providers/qwen3vl:Qwen 实现整体移出内核(git mv),实现公共 SPI 并自注册 - internal/memory/vector:新增 ProviderAdapter(公共 SPI → 内部小接口); ErrModalityUnsupported 改为公共哨兵别名;删除 VideoEmbedder 可选接口 (那正是「核心为每个新模态长方法」的坏味道) - http embedder 也变成普通 provider(注册名 http) - cmd/homed:删除 qwen import 与 onnx/http 分支,改为按 provider 名打开 + 透传 options.*;provider 打开失败只警告并禁用多模态检索,不影响启动 - config:multimodal_space.type/onnx./http.* → provider + options.* - 删除 internal/memory/qwen(整体搬迁) 测试: - pkg/embedding:注册表隔离/未知名字/非法 Info 自动关闭/ValidateVector - vector:适配器原样透传字节与 MIME、维度错被拦、Close 幂等且停止使用、 两个哨兵 errors.Is 互通 - providers/qwen3vl:新增公共 SPI 全链路集成测试(Open→Info→Embed→ 未知模态哨兵),并明确断言 Info 不声明 video 已知未完成(不得当作已验证): - 视频冻结回归 TestEmbedderVideoMatchesONNXReference **显式跳过**:Go 侧 video 模板缺少 processor 按时间组插入的字面时间戳文本 (<0.0 seconds>/<1.0 seconds>),同一输入 Python seq=1190(1152+38)、 Go 只有 22 个文本 token。时间戳也占 M-RoPE 位置,故现有 M-RoPE 自洽断言 通过不能证明与官方实现一致。修复属 provider 内部工作。 - 视觉侧三档已导出并逐档校验通过(cos 1.000000119/1.000000119/1.000000000) 验证:go build ./... ;go vet -tags onnxruntime ./... ; go test -short ./internal/memory/... ./internal/agent/core/... ./internal/sdk/... ./pkg/... ;onnxruntime 下 providers/qwen3vl 全绿(视频为显式 skip)
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 qwen3vl
|
||
|
||
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
|
||
}
|