feat(memory): 千问文本塔 ONNX 嵌入器(onnxruntime 标签,含 stub)

与 internal/memory/clip 同模式:`//go:build onnxruntime` 编真实实现,无标签时
走 stub,默认构建不链接 onnxruntime、行为不变。

加载契约(目录由 core.memory.multimodal_space.model_dir 指定):
TextTower.onnx + 外部权重分片、tokenizer.json、embed_config.json。
图内已含 last-token 池化,输出即 [batch, dim];L2 归一化在 Go 侧做。

两个刻意的设计选择:

1. **EmbedImageDense 明确报错,不返回零向量**
   导出的是文本塔,视觉塔未导出。返回零向量会让「写入了但检索不到」,
   把跨模态检索失效变成静默故障;明确报错则调用方(mediaref.go)log 后
   跳过写向量,文本路径不受影响。

2. **Fingerprint 只哈希图文件 + 配置 + 外部权重的文件名与大小**
   该目录有 6.5GB 权重分片,启动时全读一遍要几十秒、会阻塞 homeagent 启动。
   换模型必然改变文件集合或大小,足以识别切换;代价是理论上存在
   「大小相同但内容不同」的漏判,对本地单机部署可接受。已写入注释。

模板渲染(renderInstructionInput)放在无构建标签的 tokenizer.go,因此可被
测试覆盖:参考数据里有该模板串的用例,逐 token 对齐验证——模板差一个字符,
池化取到的「最后一个有效 token」位置就变,向量就不同,且不会报错。

验证:6 个测试全绿;go build ./...(stub)与 go build -tags onnxruntime
(真实实现)均通过。
This commit is contained in:
JianFeeeee
2026-09-11 00:16:45 +08:00
parent e985151da9
commit 8e88ae789f
4 changed files with 329 additions and 0 deletions

View File

@ -0,0 +1,255 @@
//go:build onnxruntime
// Package qwen 的 ONNX 推理实现(构建标签 onnxruntime与 internal/memory/clip 同模式)。
//
// 加载契约(目录由 core.memory.multimodal_space.model_dir 指定):
//
// TextTower.onnx + 外部权重分片 — 文本塔图input_ids/attention_mask → embedding
// tokenizer.json — 字节级 BPE 词表与 merges
// embed_config.json — dim / max_length / instruction / pooling
//
// 图内已含 last-token 池化,输出即 [batch, dim]L2 归一化在 Go 侧做(便宜且便于测试)。
package qwen
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"math"
"os"
"path/filepath"
"sort"
"strings"
"sync"
ort "github.com/yalue/onnxruntime_go"
)
// embedConfig 对应导出脚本产出的 embed_config.json。
type embedConfig struct {
Dimension int `json:"dim"`
MaxLength int `json:"max_length"`
Instruction string `json:"instruction"`
Pooling string `json:"pooling"`
}
// Embedder 是基于内嵌 ONNX 文本塔的稠密嵌入器。
//
// 只提供文本能力导出的是文本塔视觉塔未导出。EmbedImageDense 会明确报错,
// 而不是返回一个「看起来能用」的零向量——后者会让跨模态检索静默失效。
type Embedder struct {
mu sync.RWMutex
loaded bool
config embedConfig
tok *Tokenizer
sess *ort.DynamicAdvancedSession
fp string
}
// New 从模型目录加载文本塔。
func New(modelDir string) (*Embedder, error) {
if modelDir == "" {
return nil, fmt.Errorf("qwen model dir not specified")
}
cfgRaw, err := os.ReadFile(filepath.Join(modelDir, "embed_config.json"))
if err != nil {
return nil, fmt.Errorf("read embed_config.json: %w", err)
}
var cfg embedConfig
if err := json.Unmarshal(cfgRaw, &cfg); err != nil {
return nil, fmt.Errorf("parse embed_config.json: %w", err)
}
if cfg.Dimension <= 0 {
return nil, fmt.Errorf("embed_config.json 的 dim 无效: %d", cfg.Dimension)
}
if cfg.MaxLength <= 0 {
cfg.MaxLength = 512
}
if cfg.Pooling != "" && cfg.Pooling != "last_token" {
return nil, fmt.Errorf("不支持的池化方式 %q导出脚本只产出 last_token", cfg.Pooling)
}
tok, err := LoadTokenizer(modelDir)
if err != nil {
return nil, err
}
tok.MaxLen = cfg.MaxLength
if !ort.IsInitialized() {
if lib := findOnnxLib(); lib != "" {
ort.SetSharedLibraryPath(lib)
}
if err := ort.InitializeEnvironment(); err != nil {
return nil, fmt.Errorf("init onnx env: %w", err)
}
}
sess, err := ort.NewDynamicAdvancedSession(
filepath.Join(modelDir, "TextTower.onnx"),
[]string{"input_ids", "attention_mask"},
[]string{"embedding"},
nil,
)
if err != nil {
return nil, fmt.Errorf("create text tower session: %w", err)
}
return &Embedder{
loaded: true,
config: cfg,
tok: tok,
sess: sess,
fp: computeFingerprint(modelDir),
}, nil
}
// renderInput 按模型自带的对话模板拼输入(实现在 tokenizer.go无构建标签
func (e *Embedder) renderInput(text string) string {
return renderInstructionInput(e.config.Instruction, text)
}
// VectorizeDense 把文本编码为 L2 归一化的稠密向量。
func (e *Embedder) VectorizeDense(text string) ([]float64, error) {
e.mu.RLock()
defer e.mu.RUnlock()
if !e.loaded {
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: 分词结果为空")
}
seq := len(ids)
inputIDs := make([]int64, seq)
attn := make([]int64, seq)
for i, id := range ids {
inputIDs[i] = int64(id)
attn[i] = 1
}
idTensor, err := ort.NewTensor(ort.Shape{1, int64(seq)}, inputIDs)
if err != nil {
return nil, fmt.Errorf("input_ids tensor: %w", err)
}
defer idTensor.Destroy()
maskTensor, err := ort.NewTensor(ort.Shape{1, int64(seq)}, attn)
if err != nil {
return nil, fmt.Errorf("attention_mask tensor: %w", err)
}
defer maskTensor.Destroy()
outTensor, err := ort.NewEmptyTensor[float32](ort.Shape{1, int64(e.config.Dimension)})
if err != nil {
return nil, fmt.Errorf("output tensor: %w", err)
}
defer outTensor.Destroy()
if err := e.sess.Run([]ort.Value{idTensor, maskTensor}, []ort.Value{outTensor}); err != nil {
return nil, fmt.Errorf("text tower run: %w", err)
}
raw := outTensor.GetData()
out := make([]float64, len(raw))
var norm float64
for i, v := range raw {
out[i] = float64(v)
norm += out[i] * out[i]
}
if norm > 0 {
norm = math.Sqrt(norm)
for i := range out {
out[i] /= norm
}
}
return out, nil
}
// EmbedImageDense 不支持:导出的是**文本塔**,视觉塔未导出。
//
// 明确报错而不是返回零向量或占位调用方mediaref.go会 log 后跳过写向量,
// 若返回零向量则「写入了但检索不到」,失败会静默化。要支持图像检索需另外
// 导出视觉塔并实现 Qwen3-VL 的图像预处理patch/merge/缩放规则)。
func (e *Embedder) EmbedImageDense(_ []byte, _ string) ([]float64, error) {
return nil, fmt.Errorf("qwen text tower 不支持图像嵌入;图像检索请用 clip 或 http 路径")
}
func (e *Embedder) Fingerprint() string { return e.fp }
func (e *Embedder) Dim() int { return e.config.Dimension }
func (e *Embedder) Loaded() bool {
e.mu.RLock()
defer e.mu.RUnlock()
return e.loaded
}
func (e *Embedder) Close() {
e.mu.Lock()
defer e.mu.Unlock()
if e.sess != nil {
e.sess.Destroy()
e.sess = nil
}
e.loaded = false
}
// computeFingerprint 计算模型指纹,用于 vec_model 持久化与切换后重算判定。
//
// 为什么不直接哈希全部权重:这个模型目录有 6.5GB 外部权重分片,启动时读一遍
// 要几十秒,会阻塞 homeagent 启动。这里哈希「图文件 + 配置 + 全部外部权重的
// 文件名与大小」——换模型(哪怕只是换了权重)几乎必然改变文件集合或大小,
// 足以识别切换;代价是理论上存在「大小相同但内容不同」的漏判,对本地单机
// 部署可接受。
func computeFingerprint(modelDir string) string {
h := sha256.New()
for _, name := range []string{"TextTower.onnx", "embed_config.json"} {
if data, err := os.ReadFile(filepath.Join(modelDir, name)); err == nil {
h.Write([]byte(name))
h.Write([]byte{0})
h.Write(data)
h.Write([]byte{0})
}
}
entries, _ := os.ReadDir(modelDir)
var names []string
for _, e := range entries {
n := e.Name()
// 外部权重分片torch 新版导出器使用 onnx__<op>_<id> 与模型张量同名文件。
if strings.HasPrefix(n, "onnx__") || strings.HasSuffix(n, ".weight") {
names = append(names, n)
}
}
sort.Strings(names)
for _, n := range names {
info, err := os.Stat(filepath.Join(modelDir, n))
if err != nil {
continue
}
fmt.Fprintf(h, "%s:%d\n", n, info.Size())
}
return hex.EncodeToString(h.Sum(nil))
}
// findOnnxLib 在常见路径中查找 libonnxruntime.so。
func findOnnxLib() string {
for _, p := range []string{
"/opt/onnxruntime/lib/libonnxruntime.so",
"/usr/local/lib/libonnxruntime.so",
"/usr/lib/libonnxruntime.so",
} {
if _, err := os.Stat(p); err == nil {
return p
}
}
return ""
}

View File

@ -0,0 +1,28 @@
//go:build !onnxruntime
package qwen
import "fmt"
// Embedder 在未启用 onnxruntime 时为 no-op 实现(与 internal/memory/clip 同模式)。
// 默认构建不链接 onnxruntime保持原有 fastText/TF-IDF 行为不变。
type Embedder struct {
loaded bool
}
func New(_ string) (*Embedder, error) {
return nil, fmt.Errorf("qwen embedder requires build tag 'onnxruntime' (go build -tags onnxruntime)")
}
func (e *Embedder) Fingerprint() string { return "" }
func (e *Embedder) Dim() int { return 0 }
func (e *Embedder) Loaded() bool { return e.loaded }
func (e *Embedder) Close() {}
func (e *Embedder) VectorizeDense(_ string) ([]float64, error) {
return nil, fmt.Errorf("qwen embedder not available")
}
func (e *Embedder) EmbedImageDense(_ []byte, _ string) ([]float64, error) {
return nil, fmt.Errorf("qwen embedder not available")
}

View File

@ -136,6 +136,24 @@ func (t *Tokenizer) SpecialID(content string) (int, bool) {
return 0, false
}
// DefaultInstruction 是导出脚本随 embed_config.json 写入的默认指令。
const DefaultInstruction = "Represent the user's input."
// renderInstructionInput 按模型自带的对话模板拼输入(无构建标签,便于测试)。
//
// 必须与 HuggingFace processor 的 apply_chat_template(add_generation_prompt=True)
// 产出完全一致:指令放 system、正文放 user、以 assistant 起始符结尾。差一个
// 特殊 token池化取到的「最后一个有效 token」位置就变了嵌入也就不同——
// 而且不会报错。参考数据集里有该模板串的用例,能逐 token 对齐验证。
func renderInstructionInput(instruction, text string) string {
if instruction == "" {
instruction = DefaultInstruction
}
return "<|im_start|>system\n" + instruction +
"<|im_end|>\n<|im_start|>user\n" + text +
"<|im_end|>\n<|im_start|>assistant\n"
}
// Encode 把文本编码为 token id 序列(不含特殊 token、不做截断
func (t *Tokenizer) Encode(text string) []int {
var ids []int

View File

@ -144,6 +144,34 @@ func TestTokenizerEdgeCases(t *testing.T) {
}
}
// 模板渲染必须与参考数据里的整串完全一致,且逐 token 对齐。
//
// 这是嵌入正确性的前提:模板差一个字符,池化取到的「最后一个有效 token」
// 位置就变了,向量也就不同——而且不会报错。
func TestRenderInstructionInputMatchesTemplate(t *testing.T) {
ref := loadRef(t)
tok := loadTokenizer(t)
const want = "<|im_start|>system\nRepresent the user's input.<|im_end|>\n" +
"<|im_start|>user\n你好<|im_end|>\n<|im_start|>assistant\n"
got := renderInstructionInput("", "你好")
if got != want {
t.Fatalf("模板渲染不一致:\n got %q\n want %q", got, want)
}
for _, c := range ref.Cases {
if c.Text != want {
continue
}
if ids := tok.Encode(got); !sameIDs(ids, c.IDs) {
t.Fatalf("模板串 token 不一致:\n got %v\n want %v", ids, c.IDs)
}
return
}
t.Fatal("参考数据里缺少该模板串用例")
}
// byteEnc 必须是双射256 个字节映射到 256 个互不相同的码点。
// 有碰撞就会让不同字节编成同一个 token静默产生错误输入。
func TestBytesToUnicodeBijective(t *testing.T) {