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:
JianFeeeee
2026-09-11 03:09:01 +08:00
parent 8e88ae789f
commit 96c1d7baae
4 changed files with 134 additions and 12 deletions

View File

@ -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__<op>_<id> 与模型张量同名文件。
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",

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

View File

@ -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

View File

@ -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) {