Files
HomeAgent/internal/memory/qwen/embedder_onnx_test.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

71 lines
2.3 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.

//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())
}
}