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

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