From 96c1d7baaec092df14b172216b4bb66675b95f02 Mon Sep 17 00:00:00 2001 From: JianFeeeee Date: Fri, 11 Sep 2026 03:09:01 +0800 Subject: [PATCH] =?UTF-8?q?fix(memory):=20=E5=8D=83=E9=97=AE=E6=96=87?= =?UTF-8?q?=E6=9C=AC=E5=A1=94=E8=A1=A5=E9=BD=90=20tokenizer=20post=5Fproce?= =?UTF-8?q?ssor=20=E4=B8=8E=E7=9C=9F=E5=AE=9E=20ONNX=20=E5=9B=9E=E5=BD=92?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 验证 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。 --- internal/memory/qwen/embedder.go | 20 +++---- internal/memory/qwen/embedder_onnx_test.go | 70 ++++++++++++++++++++++ internal/memory/qwen/tokenizer.go | 24 +++++++- internal/memory/qwen/tokenizer_test.go | 32 ++++++++++ 4 files changed, 134 insertions(+), 12 deletions(-) create mode 100644 internal/memory/qwen/embedder_onnx_test.go diff --git a/internal/memory/qwen/embedder.go b/internal/memory/qwen/embedder.go index 757af6f..1a52e4a 100644 --- a/internal/memory/qwen/embedder.go +++ b/internal/memory/qwen/embedder.go @@ -2,7 +2,7 @@ // Package qwen 的 ONNX 推理实现(构建标签 onnxruntime,与 internal/memory/clip 同模式)。 // -// 加载契约(目录由 core.memory.multimodal_space.model_dir 指定): +// 加载契约:调用方传入模型目录,内核不硬编码模型名。 // // TextTower.onnx + 外部权重分片 — 文本塔图(input_ids/attention_mask → embedding) // tokenizer.json — 字节级 BPE 词表与 merges @@ -28,10 +28,10 @@ import ( // embedConfig 对应导出脚本产出的 embed_config.json。 type embedConfig struct { - Dimension int `json:"dim"` - MaxLength int `json:"max_length"` + Dimension int `json:"dim"` + MaxLength int `json:"max_length"` Instruction string `json:"instruction"` - Pooling string `json:"pooling"` + Pooling string `json:"pooling"` } // Embedder 是基于内嵌 ONNX 文本塔的稠密嵌入器。 @@ -119,12 +119,9 @@ func (e *Embedder) VectorizeDense(text string) ([]float64, error) { return nil, fmt.Errorf("qwen embedder not loaded") } - ids := e.tok.Encode(e.renderInput(text)) - if len(ids) > e.config.MaxLength { - ids = ids[:e.config.MaxLength] - } - if len(ids) == 0 { - return nil, fmt.Errorf("qwen: 分词结果为空") + ids, err := e.tok.encodeModelInput(e.renderInput(text), e.config.MaxLength) + if err != nil { + return nil, err } seq := len(ids) @@ -225,7 +222,7 @@ func computeFingerprint(modelDir string) string { for _, e := range entries { n := e.Name() // 外部权重分片:torch 新版导出器使用 onnx___ 与模型张量同名文件。 - if strings.HasPrefix(n, "onnx__") || strings.HasSuffix(n, ".weight") { + if strings.HasPrefix(n, "onnx__") || strings.HasSuffix(n, ".weight") || strings.HasSuffix(n, ".onnx.data") { names = append(names, n) } } @@ -243,6 +240,7 @@ func computeFingerprint(modelDir string) string { // findOnnxLib 在常见路径中查找 libonnxruntime.so。 func findOnnxLib() string { for _, p := range []string{ + "/opt/onnxruntime/libonnxruntime.so", "/opt/onnxruntime/lib/libonnxruntime.so", "/usr/local/lib/libonnxruntime.so", "/usr/lib/libonnxruntime.so", diff --git a/internal/memory/qwen/embedder_onnx_test.go b/internal/memory/qwen/embedder_onnx_test.go new file mode 100644 index 0000000..8d80f64 --- /dev/null +++ b/internal/memory/qwen/embedder_onnx_test.go @@ -0,0 +1,70 @@ +//go:build onnxruntime + +package qwen + +import ( + "math" + "os" + "testing" +) + +// 本地模型路径与 tokenizer_test.go 共用;产物不在仓库时跳过。 +const onnxModelDir = "/home/newqqagent/models/qwen3-vl-embed-text-onnx" + +// TestEmbedderMatchesONNXReference 冻结一条由 Python onnxruntime 1.28.0 生成的 +// FP32 参考向量,验证完整 Go 路径:模板渲染 → BPE → ONNX → L2 normalize。 +// +// 只校验前 12 维不是为了放宽正确性,而是避免把 2048 个浮点常量塞进仓库; +// tokenizer 的全部 token 已由 TestTokenizerMatchesReference 逐条精确校验,图本身 +// 另有 PyTorch↔ONNX 的多形状验证。这里负责捕获 Go 张量形状、输入名、输出名、 +// 池化/归一化或模板接线错误。 +func TestEmbedderMatchesONNXReference(t *testing.T) { + if _, err := os.Stat(onnxModelDir + "/TextTower.onnx"); err != nil { + t.Skipf("ONNX 产物不可用,跳过: %v", err) + } + + e, err := New(onnxModelDir) + if err != nil { + t.Fatalf("New: %v", err) + } + defer e.Close() + + ids, err := e.tok.encodeModelInput(e.renderInput("今天天气怎么样"), e.config.MaxLength) + if err != nil { + t.Fatalf("encodeModelInput: %v", err) + } + postID, ok := e.tok.SpecialID("<|endoftext|>") + if !ok || len(ids) != 23 || ids[len(ids)-1] != postID { + t.Fatalf("模型输入 post-processor 异常: len=%d tail=%v postID=%d ok=%v", len(ids), ids[len(ids)-1:], postID, ok) + } + + got, err := e.VectorizeDense("今天天气怎么样") + if err != nil { + t.Fatalf("VectorizeDense: %v", err) + } + if len(got) != 2048 { + t.Fatalf("向量维度 = %d,期望 2048", len(got)) + } + + want := []float64{ + -0.0288955811, 0.0522339381, 0.0360119902, -0.000676361844, + -0.0431003496, 0.00342374574, -0.000226021366, 0.0157372113, + -0.028889874, 0.0264931992, 0.0117276432, -0.00513622677, + } + for i := range want { + if diff := math.Abs(got[i] - want[i]); diff > 2e-5 { + t.Errorf("维度 %d = %.10g,参考 %.10g,差 %.3g", i, got[i], want[i], diff) + } + } + + var norm float64 + for _, v := range got { + norm += v * v + } + if diff := math.Abs(math.Sqrt(norm) - 1); diff > 1e-6 { + t.Errorf("L2 norm = %.9f,期望 1", math.Sqrt(norm)) + } + if !e.Loaded() || e.Dim() != 2048 || e.Fingerprint() == "" { + t.Errorf("元数据异常: loaded=%v dim=%d fingerprint=%q", e.Loaded(), e.Dim(), e.Fingerprint()) + } +} diff --git a/internal/memory/qwen/tokenizer.go b/internal/memory/qwen/tokenizer.go index 88d0434..8eaa027 100644 --- a/internal/memory/qwen/tokenizer.go +++ b/internal/memory/qwen/tokenizer.go @@ -154,7 +154,8 @@ func renderInstructionInput(instruction, text string) string { "<|im_end|>\n<|im_start|>assistant\n" } -// Encode 把文本编码为 token id 序列(不含特殊 token、不做截断)。 +// Encode 把文本编码为 token id 序列(识别输入中已有的特殊 token, +// 但不执行 tokenizer.json 的 post_processor,也不做截断)。 func (t *Tokenizer) Encode(text string) []int { var ids []int for _, seg := range t.splitSpecials(text) { @@ -167,6 +168,27 @@ func (t *Tokenizer) Encode(text string) []int { 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 diff --git a/internal/memory/qwen/tokenizer_test.go b/internal/memory/qwen/tokenizer_test.go index b584c98..bb9aadf 100644 --- a/internal/memory/qwen/tokenizer_test.go +++ b/internal/memory/qwen/tokenizer_test.go @@ -4,6 +4,7 @@ import ( "encoding/json" "os" "path/filepath" + "strings" "testing" ) @@ -172,6 +173,37 @@ func TestRenderInstructionInputMatchesTemplate(t *testing.T) { t.Fatal("参考数据里缺少该模板串用例") } +// 模型输入还要执行 tokenizer.json 的 TemplateProcessing:末尾追加 +// <|endoftext|>;超长输入先给正文留 maxLen-1 个位置,再保留 post token。 +func TestEncodeModelInputPostProcessor(t *testing.T) { + tok := loadTokenizer(t) + postID, ok := tok.SpecialID("<|endoftext|>") + if !ok { + t.Fatal("tokenizer 缺少 <|endoftext|>") + } + + shortRaw := tok.Encode("你好") + short, err := tok.encodeModelInput("你好", 512) + if err != nil { + t.Fatalf("短文本 encodeModelInput: %v", err) + } + if len(short) != len(shortRaw)+1 || short[len(short)-1] != postID { + t.Fatalf("短文本 post-processor 异常: raw=%v model=%v", shortRaw, short) + } + + longRaw := tok.Encode(strings.Repeat("记忆", 600)) + long, err := tok.encodeModelInput(strings.Repeat("记忆", 600), 512) + if err != nil { + t.Fatalf("长文本 encodeModelInput: %v", err) + } + if len(long) != 512 || long[511] != postID { + t.Fatalf("长文本截断异常: len=%d tail=%v", len(long), long[len(long)-1:]) + } + if !sameIDs(long[:511], longRaw[:511]) { + t.Fatal("长文本正文未按 maxLen-1 截断") + } +} + // byteEnc 必须是双射:256 个字节映射到 256 个互不相同的码点。 // 有碰撞就会让不同字节编成同一个 token,静默产生错误输入。 func TestBytesToUnicodeBijective(t *testing.T) {