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:
@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user