mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +00:00
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:
255
internal/memory/qwen/embedder.go
Normal file
255
internal/memory/qwen/embedder.go
Normal 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 ""
|
||||
}
|
||||
28
internal/memory/qwen/embedder_stub.go
Normal file
28
internal/memory/qwen/embedder_stub.go
Normal 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")
|
||||
}
|
||||
@ -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
|
||||
|
||||
@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user