mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-04 00:03:59 +00:00
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。
This commit is contained in:
70
internal/memory/qwen/embedder_onnx_test.go
Normal file
70
internal/memory/qwen/embedder_onnx_test.go
Normal file
@ -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())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user