From 8e88ae789f6fad8b1bb836ede0a0d01e45b08b68 Mon Sep 17 00:00:00 2001 From: JianFeeeee Date: Fri, 11 Sep 2026 00:16:45 +0800 Subject: [PATCH] =?UTF-8?q?feat(memory):=20=E5=8D=83=E9=97=AE=E6=96=87?= =?UTF-8?q?=E6=9C=AC=E5=A1=94=20ONNX=20=E5=B5=8C=E5=85=A5=E5=99=A8?= =?UTF-8?q?=EF=BC=88onnxruntime=20=E6=A0=87=E7=AD=BE=EF=BC=8C=E5=90=AB=20s?= =?UTF-8?q?tub=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 与 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 (真实实现)均通过。 --- internal/memory/qwen/embedder.go | 255 +++++++++++++++++++++++++ internal/memory/qwen/embedder_stub.go | 28 +++ internal/memory/qwen/tokenizer.go | 18 ++ internal/memory/qwen/tokenizer_test.go | 28 +++ 4 files changed, 329 insertions(+) create mode 100644 internal/memory/qwen/embedder.go create mode 100644 internal/memory/qwen/embedder_stub.go diff --git a/internal/memory/qwen/embedder.go b/internal/memory/qwen/embedder.go new file mode 100644 index 0000000..757af6f --- /dev/null +++ b/internal/memory/qwen/embedder.go @@ -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___ 与模型张量同名文件。 + 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 "" +} diff --git a/internal/memory/qwen/embedder_stub.go b/internal/memory/qwen/embedder_stub.go new file mode 100644 index 0000000..aba034c --- /dev/null +++ b/internal/memory/qwen/embedder_stub.go @@ -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") +} diff --git a/internal/memory/qwen/tokenizer.go b/internal/memory/qwen/tokenizer.go index 213e871..88d0434 100644 --- a/internal/memory/qwen/tokenizer.go +++ b/internal/memory/qwen/tokenizer.go @@ -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 diff --git a/internal/memory/qwen/tokenizer_test.go b/internal/memory/qwen/tokenizer_test.go index 41639bf..b584c98 100644 --- a/internal/memory/qwen/tokenizer_test.go +++ b/internal/memory/qwen/tokenizer_test.go @@ -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) {