Files
HomeAgent/internal/memory/clip/embedder.go
JianFeeeee 6f8056d236 refactor(clip): CLIP 收敛为稠密 MultimodalEmbedder,不污染稀疏 Vectorizer/TF-IDF 语义
文本相似度检索是层次化系统:TF-IDF 高频削弱加权(idf<0.1 丢弃)+
倒排剪枝(只召回共享特征者)+ cosine。CLIP 512 维稠密向量若以
map[string]float64 稀疏形式实现 vector.Vectorizer 并塞进 vector.Store,
会让 512 维全部成为倒排 key → 候选集≈全库、剪枝失效,且绕过 TF-IDF
高频削弱,与既有文本检索语义错配。

收敛:
- vector.MultimodalEmbedder 改为独立稠密接口(VectorizeDense/
  EmbedImageDense/Fingerprint/Dim/Loaded/Close),不再继承稀疏 Vectorizer
- clip.Embedder 删除稀疏垫片 Vectorize/EmbedImage/denseToVector,
  只产出稠密向量;文档/知识/上下文层继续用 TF-IDF/fastText 稀疏路径
- 分层明确:文本→文本走 TF-IDF/fastText;文本↔图像、图像↔图像走
  CLIP 稠密 QueryMedia(媒体层独立稠密余弦,原样保留)
2026-09-09 10:26:10 +08:00

485 lines
13 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

//go:build onnxruntime
// Package clip 提供基于 CLIP ONNX 的稠密多模态编码器。
//
// 构建标签 onnxruntime 控制是否编译此实现(与 internal/nlp/onnx.go 同模式)。
// 未配置 clip_model_dir 时不会初始化 ONNX Runtime现有 fastText/TF-IDF 行为不变。
//
// 支持的模型文件(统一放置于 clip_model_dir 目录):
//
// text.onnx — CLIP 文本编码器input_ids + attention_mask → text_embed [1,512]
// vision.onnx — CLIP 图像编码器pixel_values → image_embed [1,512]
// clip_config.json — 模型元数据dimension, context_length, image_size, mean, std
// tokenizer.json — HuggingFace tokenizer.json含 vocab + merges
// merges.txt — BPE merges 文件
//
// 设计:只产出**稠密** 512 维向量VectorizeDense / EmbedImageDense供媒体层
// media.Store 的稠密余弦检索QueryMedia消费。刻意不实现 vector.Vectorizer
// (稀疏词向量)——文档/知识/上下文层的文本相似度检索保留 TF-IDF 高频削弱加权 +
// 倒排剪枝 + fastText 稀疏词向量的既有设计CLIP 稠密空间不混入那套索引。
package clip
import (
"bytes"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"image"
_ "image/jpeg"
_ "image/png"
"log"
"math"
"os"
"path/filepath"
"strings"
"sync"
ort "github.com/yalue/onnxruntime_go"
)
// clipConfig 描述模型的超参数与归一化常数。
type clipConfig struct {
Model string `json:"model"`
Dimension int `json:"dimension"`
ContextLength int `json:"context_length"`
ImageSize int `json:"image_size"`
Mean []float64 `json:"mean"`
Std []float64 `json:"std"`
}
// Embedder 是 CLIP 多模态稠密编码器,实现 vector.MultimodalEmbedder。
// 产出 512 维共享稠密空间向量,供 media.Store 的稠密余弦检索。
//
// 刻意**不**实现 vector.Vectorizer稀疏词向量CLIP 稠密向量若以
// map[string]float64 稀疏形式塞进文档/知识层的 vector.Store会破坏其
// TF-IDF 高频削弱加权与倒排剪枝语义。文本层的相似度检索保持 TF-IDF/fastText。
type Embedder struct {
mu sync.RWMutex
config clipConfig
vocab map[string]int64
merges []string
textSess *ort.DynamicAdvancedSession
imgSess *ort.DynamicAdvancedSession
close sync.Once
loaded bool
fingerprint string
}
// Fingerprint 返回当前模型目录的指纹(文本+视觉模型文件 SHA256 拼接),
// 用于检测模型切换后触发重算。
func (e *Embedder) Fingerprint() string {
e.mu.RLock()
defer e.mu.RUnlock()
return e.fingerprint
}
// Dim 返回向量维度。
func (e *Embedder) Dim() int {
e.mu.RLock()
defer e.mu.RUnlock()
return e.config.Dimension
}
// Loaded 返回加载状态。
func (e *Embedder) Loaded() bool {
e.mu.RLock()
defer e.mu.RUnlock()
return e.loaded
}
// VectorizeDense 将文本转为归一化的 []float64 向量CLIP 共享空间)。
func (e *Embedder) VectorizeDense(text string) ([]float64, error) {
e.mu.RLock()
defer e.mu.RUnlock()
if !e.loaded {
return nil, fmt.Errorf("clip embedder not loaded")
}
tokens := tokenizeCLIP(text, e.vocab, e.merges, e.config.ContextLength)
if len(tokens) == 0 {
return make([]float64, e.config.Dimension), nil
}
dim := e.config.Dimension
inputIDs := make([]int64, e.config.ContextLength)
attnMask := make([]int64, e.config.ContextLength)
for i, tok := range tokens {
if i >= e.config.ContextLength {
break
}
inputIDs[i] = tok
attnMask[i] = 1
}
idTensor, err := ort.NewTensor(ort.Shape{1, int64(e.config.ContextLength)}, inputIDs)
if err != nil {
return nil, fmt.Errorf("input_ids tensor: %w", err)
}
defer idTensor.Destroy()
maskTensor, err := ort.NewTensor(ort.Shape{1, int64(e.config.ContextLength)}, attnMask)
if err != nil {
return nil, fmt.Errorf("attention_mask tensor: %w", err)
}
defer maskTensor.Destroy()
featTensor, err := ort.NewEmptyTensor[float32](ort.Shape{1, int64(dim)})
if err != nil {
return nil, fmt.Errorf("output tensor: %w", err)
}
defer featTensor.Destroy()
if err := e.textSess.Run([]ort.Value{idTensor, maskTensor}, []ort.Value{featTensor}); err != nil {
return nil, fmt.Errorf("text run: %w", err)
}
raw := featTensor.GetData()
out := make([]float64, dim)
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 将图像字节转为归一化的 []float64 向量CLIP 共享空间)。
func (e *Embedder) EmbedImageDense(img []byte, mime string) ([]float64, error) {
e.mu.RLock()
defer e.mu.RUnlock()
if !e.loaded {
return nil, fmt.Errorf("clip embedder not loaded")
}
return e.embedImageDenseUnlocked(img, mime)
}
func (e *Embedder) embedImageDenseUnlocked(img []byte, mime string) ([]float64, error) {
decoded, _, err := image.Decode(bytes.NewReader(img))
if err != nil {
return nil, fmt.Errorf("decode image: %w", err)
}
size := e.config.ImageSize
resized := resizeImage(decoded, size, size)
pixels := make([]float32, 3*size*size)
for y := 0; y < size; y++ {
for x := 0; x < size; x++ {
r, g, b, _ := resized.At(x, y).RGBA()
rf := float64(r) / 65535.0
gf := float64(g) / 65535.0
bf := float64(b) / 65535.0
for c, v := range []float64{rf, gf, bf} {
norm := (v - e.config.Mean[c]) / e.config.Std[c]
pixels[c*size*size+y*size+x] = float32(norm)
}
}
}
pixelTensor, err := ort.NewTensor(ort.Shape{1, 3, int64(size), int64(size)}, pixels)
if err != nil {
return nil, fmt.Errorf("pixel_values tensor: %w", err)
}
defer pixelTensor.Destroy()
dim := e.config.Dimension
featTensor, err := ort.NewEmptyTensor[float32](ort.Shape{1, int64(dim)})
if err != nil {
return nil, fmt.Errorf("output tensor: %w", err)
}
defer featTensor.Destroy()
if err := e.imgSess.Run([]ort.Value{pixelTensor}, []ort.Value{featTensor}); err != nil {
return nil, fmt.Errorf("vision run: %w", err)
}
raw := featTensor.GetData()
out := make([]float64, dim)
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
}
// Close 释放 ONNX Runtime 资源。
func (e *Embedder) Close() {
e.close.Do(func() {
e.mu.Lock()
defer e.mu.Unlock()
if e.textSess != nil {
e.textSess.Destroy()
}
if e.imgSess != nil {
e.imgSess.Destroy()
}
e.loaded = false
})
}
// New 从目录加载 CLIP 模型。目录需包含 text.onnx、vision.onnx、
// clip_config.json、tokenizer.json、merges.txt。
func New(modelDir string) (*Embedder, error) {
if modelDir == "" {
return nil, fmt.Errorf("clip model dir not specified")
}
// 读取配置
cfgData, err := os.ReadFile(filepath.Join(modelDir, "clip_config.json"))
if err != nil {
return nil, fmt.Errorf("read clip_config.json: %w", err)
}
var cfg clipConfig
if err := json.Unmarshal(cfgData, &cfg); err != nil {
return nil, fmt.Errorf("parse clip_config.json: %w", err)
}
if cfg.Dimension <= 0 || cfg.ContextLength <= 0 || cfg.ImageSize <= 0 {
return nil, fmt.Errorf("invalid clip config: dim=%d ctx=%d img=%d", cfg.Dimension, cfg.ContextLength, cfg.ImageSize)
}
if len(cfg.Mean) != 3 || len(cfg.Std) != 3 {
return nil, fmt.Errorf("clip config mean/std must have 3 channels")
}
// 加载 tokenizer
vocab, err := loadTokenizerVocab(filepath.Join(modelDir, "tokenizer.json"))
if err != nil {
return nil, fmt.Errorf("load tokenizer: %w", err)
}
merges, err := loadMerges(filepath.Join(modelDir, "merges.txt"))
if err != nil {
return nil, fmt.Errorf("load merges: %w", err)
}
// 初始化 ONNX Runtime只初始化一次
if !ort.IsInitialized() {
// 尝试从 nlp 同样的路径查找 libonnxruntime.so
libPath := findOnnxLib()
if libPath != "" {
ort.SetSharedLibraryPath(libPath)
}
if err := ort.InitializeEnvironment(); err != nil {
return nil, fmt.Errorf("init onnx env: %w", err)
}
}
// 创建文本编码器会话
textSess, err := ort.NewDynamicAdvancedSession(
filepath.Join(modelDir, "text.onnx"),
[]string{"input_ids", "attention_mask"},
[]string{"text_embed"},
nil,
)
if err != nil {
return nil, fmt.Errorf("create text session: %w", err)
}
// 创建视觉编码器会话
imgSess, err := ort.NewDynamicAdvancedSession(
filepath.Join(modelDir, "vision.onnx"),
[]string{"pixel_values"},
[]string{"image_embed"},
nil,
)
if err != nil {
textSess.Destroy()
return nil, fmt.Errorf("create vision session: %w", err)
}
// 计算模型指纹
fp := computeFingerprint(modelDir)
log.Printf("[clip] loaded %s dim=%d ctx=%d img=%d from %s (fp=%s)", cfg.Model, cfg.Dimension, cfg.ContextLength, cfg.ImageSize, modelDir, fp[:12])
return &Embedder{
config: cfg,
vocab: vocab,
merges: merges,
textSess: textSess,
imgSess: imgSess,
loaded: true,
fingerprint: fp,
}, nil
}
// computeFingerprint 计算模型文件指纹text.onnx + vision.onnx 的 SHA256
func computeFingerprint(modelDir string) string {
h := sha256.New()
for _, name := range []string{"text.onnx", "vision.onnx"} {
data, err := os.ReadFile(filepath.Join(modelDir, name))
if err != nil {
continue
}
h.Write(data)
h.Write([]byte{0}) // 分隔符
}
return hex.EncodeToString(h.Sum(nil))
}
// findOnnxLib 在常见路径中查找 libonnxruntime.so。
func findOnnxLib() string {
for _, p := range []string{
"/opt/onnxruntime/libonnxruntime.so",
"libonnxruntime.so",
} {
if _, err := os.Stat(p); err == nil {
return p
}
}
return ""
}
// resizeImage 使用最近邻将 src 缩放到 dstW×dstH。
// 生产中应使用双线性插值,此处为 MVP 简化。
func resizeImage(src image.Image, dstW, dstH int) image.Image {
srcB := src.Bounds()
srcW := srcB.Dx()
srcH := srcB.Dy()
if srcW == dstW && srcH == dstH {
return src
}
dst := image.NewRGBA(image.Rect(0, 0, dstW, dstH))
for y := 0; y < dstH; y++ {
for x := 0; x < dstW; x++ {
sx := srcB.Min.X + x*srcW/dstW
sy := srcB.Min.Y + y*srcH/dstH
dst.Set(x, y, src.At(sx, sy))
}
}
return dst
}
// ---- BPE Tokenizer ----
// loadTokenizerVocab 从 HuggingFace tokenizer.json 中提取 vocabtoken→id 映射)。
func loadTokenizerVocab(path string) (map[string]int64, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, err
}
var tok struct {
Model struct {
Vocab map[string]int64 `json:"vocab"`
} `json:"model"`
}
if err := json.Unmarshal(data, &tok); err != nil {
return nil, err
}
if len(tok.Model.Vocab) == 0 {
return nil, fmt.Errorf("empty vocab in %s", path)
}
return tok.Model.Vocab, nil
}
// loadMerges 从 merges.txt 加载 BPE 合并规则。
func loadMerges(path string) ([]string, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, err
}
lines := strings.Split(strings.TrimSpace(string(data)), "\n")
// 第一行是版本号("#version: 0.2"),跳过
var merges []string
for _, line := range lines[1:] {
line = strings.TrimSpace(line)
if line == "" {
continue
}
merges = append(merges, line)
}
return merges, nil
}
// tokenizeCLIP 将文本分词为模型 vocab 中的 token id 序列。
//
// 此模型transformers 5.x 导出的 CLIP tokenizer.json是**词级 BPE**
// 词末 token 带 </w> 后缀("a</w>"=320、"red</w>"=736词中片段不带。
// 流程lowercase → 按空白/标点拆词 → 每词做字符级 BPE 合并 →
// 末尾片段加 </w> 查 vocab其余片段直接查查不到则丢弃。
func tokenizeCLIP(text string, vocab map[string]int64, merges []string, maxLen int) []int64 {
rank := make(map[string]int, len(merges))
for i, m := range merges {
rank[m] = i
}
const endTok = "</w>"
var tokens []int64
if id, ok := vocab["<|startoftext|>"]; ok {
tokens = append(tokens, id)
}
for _, word := range strings.Fields(strings.ToLower(text)) {
seq := make([]string, 0, len(word))
for _, ch := range word {
seq = append(seq, string(ch))
}
merged := bpeMerge(seq, rank)
for i, t := range merged {
lookup := t
if i == len(merged)-1 {
// 词末片段带 </w>
lookup = t + endTok
}
if id, ok := vocab[lookup]; ok {
tokens = append(tokens, id)
}
}
}
if id, ok := vocab["<|endoftext|>"]; ok {
tokens = append(tokens, id)
}
if len(tokens) > maxLen {
tokens = tokens[:maxLen]
}
return tokens
}
// bpeMerge 对单个词的字符序列应用 BPE 合并直到无可合并对。
// rank[pair] 越小越优先merges.txt 顺序)。
func bpeMerge(seq []string, rank map[string]int) []string {
for len(seq) > 1 {
// 找 rank 最低的可合并相邻对
bestRank := -1
bestPair := ""
for i := 0; i < len(seq)-1; i++ {
pair := seq[i] + " " + seq[i+1]
if r, ok := rank[pair]; ok && (bestRank < 0 || r < bestRank) {
bestRank = r
bestPair = pair
}
}
if bestPair == "" {
break
}
parts := strings.SplitN(bestPair, " ", 2)
merged := parts[0] + parts[1]
// 一次性合并所有相邻的该 pair
var out []string
for i := 0; i < len(seq); i++ {
if i < len(seq)-1 && seq[i] == parts[0] && seq[i+1] == parts[1] {
out = append(out, merged)
i++
} else {
out = append(out, seq[i])
}
}
seq = out
}
return seq
}