mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 15:53:56 +00:00
feat(clip): 多模态向量器(CLIP ONNX)——文本/图像 512 维共享空间 + 媒体向量写入与重算
- internal/memory/clip:CLIP ONNX 向量器(onnxruntime 构建标签控制,默认构建不链接 ONNX) - clip.New(modelDir) 加载 text.onnx/vision.onnx(输出 text_embed/image_embed [batch,512]) - 实现 vector.Vectorizer + vector.MultimodalEmbedder(Vectorize/EmbedImage + Dense 变体) - 词级 BPE tokenizer:merges 合并后词末片段带 </w> 查 vocab,与官方 encode 逐 id 对齐 - EmbedImage:解码→resize 224→NCHW→normalize→vision session - Fingerprint(text+vision 文件 sha256)供模型切换检测 - stub 版(无 onnxruntime 标签)保持默认构建行为不变 - vector/store.go:新增 MultimodalEmbedder 接口 - media.Store:新增 StaleVecDigests(currentModel)——查 vec_model 不匹配/缺失的图片 - agent core:AgentConfig.ClipEmbedder + Agent.clipEmb 接线; describePendingMedia 描述成功后 EmbedImageDense→SetVec; 新增 reembedStaleMedia 启动补算历史无向量图片 - config:core.memory.media.clip_model_dir(未配置退化为现有 fastText/TF-IDF 行为) - cmd/homed:读 clip_model_dir 加载 CLIP,失败仅记日志不阻塞启动 测试:TestSmokeLoadAndEncode(文本语义 cat>dog 0.914>physics 0.740)、 TestCrossModalAlignment(red-image vs red-text 0.063>blue -0.009,与 Python 一致)、 TestTokEnd(与官方 encode 逐 id 对齐)、TestStaleVecDigests,含 -race 全绿
This commit is contained in:
@ -24,6 +24,7 @@ import (
|
|||||||
logpkg "gitcode.com/JianFeeeee/HomeAgent/internal/log"
|
logpkg "gitcode.com/JianFeeeee/HomeAgent/internal/log"
|
||||||
luapkg "gitcode.com/JianFeeeee/HomeAgent/internal/lua"
|
luapkg "gitcode.com/JianFeeeee/HomeAgent/internal/lua"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||||
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/clip"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/media"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/media"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/pipeline"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/pipeline"
|
||||||
@ -346,6 +347,21 @@ func main() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 多模态嵌入(CLIP ONNX,可选):配置 clip_model_dir 时启用,
|
||||||
|
// 图片入库/描述时计算视觉向量,供跨模态检索;未配置则退回纯 fastText 文本路径。
|
||||||
|
var clipEmbedder *clip.Embedder
|
||||||
|
if clipDir := cfgReg.GetString("core.memory.media.clip_model_dir", ""); clipDir != "" {
|
||||||
|
e, err := clip.New(clipDir)
|
||||||
|
if err != nil {
|
||||||
|
// 配置了但加载失败:记日志降级,不阻塞启动(媒体记忆是增强项)
|
||||||
|
log.Printf("[homed] warning: clip embedder load failed: %v(多模态向量检索已禁用)", err)
|
||||||
|
} else {
|
||||||
|
clipEmbedder = e
|
||||||
|
defer e.Close()
|
||||||
|
log.Printf("[homed] clip embedder active: dim=%d fp=%s", e.Dim(), e.Fingerprint()[:min(12, len(e.Fingerprint()))])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
ks := knowledge.NewStore(filepath.Join(cfg.Daemon.DataDir, "knowledge"))
|
ks := knowledge.NewStore(filepath.Join(cfg.Daemon.DataDir, "knowledge"))
|
||||||
if err := ks.Start(); err != nil {
|
if err := ks.Start(); err != nil {
|
||||||
log.Printf("[homed] warning: knowledge store: %v", err)
|
log.Printf("[homed] warning: knowledge store: %v", err)
|
||||||
@ -468,6 +484,7 @@ func main() {
|
|||||||
ContextSavePath: filepath.Join(cfg.Daemon.DataDir, "memory", "context.json"),
|
ContextSavePath: filepath.Join(cfg.Daemon.DataDir, "memory", "context.json"),
|
||||||
EmbeddingModelPath: cfgReg.GetString("core.agent.embedding_model_path", ""),
|
EmbeddingModelPath: cfgReg.GetString("core.agent.embedding_model_path", ""),
|
||||||
Embedder: embedder,
|
Embedder: embedder,
|
||||||
|
ClipEmbedder: clipEmbedder,
|
||||||
StageHost: stageHost,
|
StageHost: stageHost,
|
||||||
EventBus: evBus,
|
EventBus: evBus,
|
||||||
ThinkingEnabled: cfg.LLM.ThinkingEnabled,
|
ThinkingEnabled: cfg.LLM.ThinkingEnabled,
|
||||||
|
|||||||
@ -13,6 +13,7 @@ import (
|
|||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/events"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/events"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/knowledge"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/knowledge"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||||
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/clip"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/media"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/media"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/social"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/social"
|
||||||
@ -144,6 +145,9 @@ type Agent struct {
|
|||||||
// 词嵌入模型,用于实体语义相似度计算
|
// 词嵌入模型,用于实体语义相似度计算
|
||||||
embedder *memory.StaticEmbedder
|
embedder *memory.StaticEmbedder
|
||||||
|
|
||||||
|
// clipEmb 是 CLIP 多模态嵌入(可选),nil 时跳过视觉向量计算。
|
||||||
|
clipEmb *clip.Embedder
|
||||||
|
|
||||||
// 技能索引提供者:由 skillmgr 插件实现,向 system prompt 注入轻量技能索引
|
// 技能索引提供者:由 skillmgr 插件实现,向 system prompt 注入轻量技能索引
|
||||||
skillIndex SkillIndexProvider
|
skillIndex SkillIndexProvider
|
||||||
}
|
}
|
||||||
@ -174,6 +178,7 @@ type AgentConfig struct {
|
|||||||
MediaGCInterval time.Duration
|
MediaGCInterval time.Duration
|
||||||
MediaGCMinAge time.Duration
|
MediaGCMinAge time.Duration
|
||||||
MediaDescribe bool
|
MediaDescribe bool
|
||||||
|
ClipEmbedder *clip.Embedder
|
||||||
Personality *agentPkg.Personality
|
Personality *agentPkg.Personality
|
||||||
PluginReg *plugin.Registry
|
PluginReg *plugin.Registry
|
||||||
PluginDir string
|
PluginDir string
|
||||||
@ -279,6 +284,7 @@ func New(cfg AgentConfig) *Agent {
|
|||||||
thinkingEnabled: cfg.ThinkingEnabled,
|
thinkingEnabled: cfg.ThinkingEnabled,
|
||||||
inputCfg: cfg.InputProcessing,
|
inputCfg: cfg.InputProcessing,
|
||||||
embedder: embedder,
|
embedder: embedder,
|
||||||
|
clipEmb: cfg.ClipEmbedder,
|
||||||
noMergeMarkers: make(map[string]int),
|
noMergeMarkers: make(map[string]int),
|
||||||
lastInput: make(map[string]time.Time),
|
lastInput: make(map[string]time.Time),
|
||||||
}
|
}
|
||||||
@ -296,6 +302,7 @@ func (a *Agent) Start() {
|
|||||||
go a.reviewLoop()
|
go a.reviewLoop()
|
||||||
go a.mediaGCLoop()
|
go a.mediaGCLoop()
|
||||||
go a.mediaDescribeLoop()
|
go a.mediaDescribeLoop()
|
||||||
|
a.reembedStaleMedia()
|
||||||
log.Printf("[agent] %s started, waiting for IO interrupts", a.id)
|
log.Printf("[agent] %s started, waiting for IO interrupts", a.id)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@ -179,5 +179,64 @@ func (a *Agent) describePendingMedia() {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
log.Printf("[media] 已描述 %s (%s, %d 字, 源=%s)", shortDigest(it.Digest), kind, len([]rune(desc)), srcName)
|
log.Printf("[media] 已描述 %s (%s, %d 字, 源=%s)", shortDigest(it.Digest), kind, len([]rune(desc)), srcName)
|
||||||
|
|
||||||
|
// 描述成功后,若 CLIP 嵌入器可用且是图片,计算视觉向量。
|
||||||
|
// 这是"描述 + 向量"两步同步完成的路径;对于历史已有描述但无向量的媒体,
|
||||||
|
// 由启动时的 reembedStaleMedia 补算。
|
||||||
|
if a.clipEmb != nil && kind == "image" && it.Kind == media.KindImage {
|
||||||
|
vec, err := a.clipEmb.EmbedImageDense(data, mime)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("[media] 视觉嵌入失败 %s: %v", shortDigest(it.Digest), err)
|
||||||
|
} else if err := a.mediaStore.SetVec(it.Digest, vec, a.clipEmb.Fingerprint()); err != nil {
|
||||||
|
log.Printf("[media] 写向量失败 %s: %v", shortDigest(it.Digest), err)
|
||||||
|
} else {
|
||||||
|
log.Printf("[media] 已嵌入 %s (dim=%d)", shortDigest(it.Digest), len(vec))
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// reembedStaleMedia 在启动时为历史已有描述但无 CLIP 向量的图片补算视觉向量。
|
||||||
|
// 避免安装 CLIP 后,旧图片永远只有描述文本、没有视觉向量,直到下次 Describe 才能写入。
|
||||||
|
func (a *Agent) reembedStaleMedia() {
|
||||||
|
if a.clipEmb == nil || a.mediaStore == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fp := a.clipEmb.Fingerprint()
|
||||||
|
digests, err := a.mediaStore.StaleVecDigests(fp)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("[media] 查询需重算向量的媒体失败: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if len(digests) == 0 {
|
||||||
|
log.Printf("[media] 无历史媒体需要补算视觉向量")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Printf("[media] 启动补算视觉向量: %d 条 (fp=%s...)", len(digests), fp[:min(12, len(fp))])
|
||||||
|
done := 0
|
||||||
|
for _, d := range digests {
|
||||||
|
it, err := a.mediaStore.Stat(d)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
data, err := a.mediaStore.Get(d)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
mime := it.MIME
|
||||||
|
if mime == "" {
|
||||||
|
mime = "image/png"
|
||||||
|
}
|
||||||
|
vec, err := a.clipEmb.EmbedImageDense(data, mime)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("[media] 启动补算失败 %s: %v", shortDigest(d), err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := a.mediaStore.SetVec(d, vec, fp); err != nil {
|
||||||
|
log.Printf("[media] 启动写入向量失败 %s: %v", shortDigest(d), err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
done++
|
||||||
|
}
|
||||||
|
log.Printf("[media] 启动补算完成: %d/%d", done, len(digests))
|
||||||
|
}
|
||||||
|
|||||||
@ -655,6 +655,7 @@ func (r *ConfigRegistry) seedCoreDefs(dataDir string) {
|
|||||||
reg(ConfigDef{Key: "core.memory.media.gc_interval", Default: "6h", Type: "duration", DisplayName: "媒体 GC 间隔", Description: "清理无引用媒体的周期;0 表示不自动清理", Category: "memory"})
|
reg(ConfigDef{Key: "core.memory.media.gc_interval", Default: "6h", Type: "duration", DisplayName: "媒体 GC 间隔", Description: "清理无引用媒体的周期;0 表示不自动清理", Category: "memory"})
|
||||||
reg(ConfigDef{Key: "core.memory.media.gc_min_age", Default: "1h", Type: "duration", DisplayName: "媒体 GC 保护期", Description: "新入库媒体在此时长内不被清理。刚落盘还没来得及挂到记忆上的项引用计数也是 0,靠这个保护期避免被误删", Category: "memory"})
|
reg(ConfigDef{Key: "core.memory.media.gc_min_age", Default: "1h", Type: "duration", DisplayName: "媒体 GC 保护期", Description: "新入库媒体在此时长内不被清理。刚落盘还没来得及挂到记忆上的项引用计数也是 0,靠这个保护期避免被误删", Category: "memory"})
|
||||||
reg(ConfigDef{Key: "core.memory.media.describe_on_ingest", Default: "false", Type: "bool", DisplayName: "自动描述媒体", Description: "后台用视觉/音频模型给未描述的媒体生成文字描述。**描述文本才是持久语义记忆**——blob 会被容量 GC 淘汰,描述会随记忆各层一直留存并可检索。代价是消耗视觉模型配额(单张图实测约 10s),故默认关闭;开启后每 30s 最多处理 4 条,不跟对话抢额度", Category: "memory"})
|
reg(ConfigDef{Key: "core.memory.media.describe_on_ingest", Default: "false", Type: "bool", DisplayName: "自动描述媒体", Description: "后台用视觉/音频模型给未描述的媒体生成文字描述。**描述文本才是持久语义记忆**——blob 会被容量 GC 淘汰,描述会随记忆各层一直留存并可检索。代价是消耗视觉模型配额(单张图实测约 10s),故默认关闭;开启后每 30s 最多处理 4 条,不跟对话抢额度", Category: "memory"})
|
||||||
|
reg(ConfigDef{Key: "core.memory.media.clip_model_dir", Default: "", Type: "string", DisplayName: "CLIP 模型目录", Description: "多模态嵌入的 CLIP ONNX 模型目录(含 text.onnx、vision.onnx、clip_config.json、tokenizer.json、merges.txt)。留空禁用多模态向量检索,只保留 fastText 文本路径;配置后图片入库时自动计算视觉向量并与文档混合检索。修改后需重启生效。", Category: "memory"})
|
||||||
reg(ConfigDef{Key: "core.knowledge.path", Default: filepath.Join(dataDir, "knowledge"), Type: "string", DisplayName: "知识库路径", Description: "知识库存储目录", Category: "paths"})
|
reg(ConfigDef{Key: "core.knowledge.path", Default: filepath.Join(dataDir, "knowledge"), Type: "string", DisplayName: "知识库路径", Description: "知识库存储目录", Category: "paths"})
|
||||||
reg(ConfigDef{Key: "core.log.path", Default: filepath.Join(dataDir, "log"), Type: "string", DisplayName: "日志目录", Description: "日志文件输出目录", Category: "paths"})
|
reg(ConfigDef{Key: "core.log.path", Default: filepath.Join(dataDir, "log"), Type: "string", DisplayName: "日志目录", Description: "日志文件输出目录", Category: "paths"})
|
||||||
|
|
||||||
|
|||||||
509
internal/memory/clip/embedder.go
Normal file
509
internal/memory/clip/embedder.go
Normal file
@ -0,0 +1,509 @@
|
|||||||
|
//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_features [1,512])
|
||||||
|
// vision.onnx — CLIP 图像编码器(pixel_values → image_features [1,512])
|
||||||
|
// clip_config.json — 模型元数据(dimension, context_length, image_size, mean, std)
|
||||||
|
// tokenizer.json — HuggingFace tokenizer.json(含 vocab + merges)
|
||||||
|
// merges.txt — BPE merges 文件(CLIP 使用的字节级 BPE)
|
||||||
|
//
|
||||||
|
// 设计:同时满足 vector.Vectorizer 接口(稀疏 map,供文档检索复用)和直接返回
|
||||||
|
// []float64 的方法(供媒体嵌入与 QueryMedia 直接调用)。
|
||||||
|
package clip
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"image"
|
||||||
|
_ "image/jpeg"
|
||||||
|
_ "image/png"
|
||||||
|
"log"
|
||||||
|
"math"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||||
|
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 实现 vector.Vectorizer,提供文本向量化与图像向量化。
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// Vectorize 将文本转为向量,供 vector.Vectorizer 接口使用(稀疏 map)。
|
||||||
|
func (e *Embedder) Vectorize(text string) vector.Vector {
|
||||||
|
dense, err := e.VectorizeDense(text)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("[clip] Vectorize 失败: %v", err)
|
||||||
|
return vector.Vector{}
|
||||||
|
}
|
||||||
|
return denseToVector(dense)
|
||||||
|
}
|
||||||
|
|
||||||
|
// EmbedImage 将图像字节转为向量,供 vector.Vectorizer 接口使用(稀疏 map)。
|
||||||
|
func (e *Embedder) EmbedImage(img []byte, mime string) (vector.Vector, error) {
|
||||||
|
dense, err := e.EmbedImageDense(img, mime)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return denseToVector(dense), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// denseToVector 将 []float64 稀疏化为 vector.Vector(CLIP 维度通常 512,不会太大)。
|
||||||
|
func denseToVector(d []float64) vector.Vector {
|
||||||
|
vec := make(vector.Vector, len(d))
|
||||||
|
for i, v := range d {
|
||||||
|
if v != 0 {
|
||||||
|
vec[strconv.Itoa(i)] = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return vec
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- BPE Tokenizer ----
|
||||||
|
|
||||||
|
// loadTokenizerVocab 从 HuggingFace tokenizer.json 中提取 vocab(token→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
|
||||||
|
}
|
||||||
34
internal/memory/clip/embedder_stub.go
Normal file
34
internal/memory/clip/embedder_stub.go
Normal file
@ -0,0 +1,34 @@
|
|||||||
|
//go:build !onnxruntime
|
||||||
|
|
||||||
|
package clip
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Embedder 在未启用 onnxruntime 时为 no-op 实现。
|
||||||
|
// 构建时不链接 onnxruntime,default 构建保持原有 fastText/TF-IDF 行为不变。
|
||||||
|
type Embedder struct {
|
||||||
|
loaded bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func New(_ string) (*Embedder, error) {
|
||||||
|
return nil, fmt.Errorf("clip 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) Vectorize(_ string) vector.Vector { return nil }
|
||||||
|
func (e *Embedder) EmbedImage(_ []byte, _ string) (vector.Vector, error) {
|
||||||
|
return nil, fmt.Errorf("clip embedder not available")
|
||||||
|
}
|
||||||
|
func (e *Embedder) VectorizeDense(_ string) ([]float64, error) {
|
||||||
|
return nil, fmt.Errorf("clip embedder not available")
|
||||||
|
}
|
||||||
|
func (e *Embedder) EmbedImageDense(_ []byte, _ string) ([]float64, error) {
|
||||||
|
return nil, fmt.Errorf("clip embedder not available")
|
||||||
|
}
|
||||||
|
func (e *Embedder) Close() {}
|
||||||
154
internal/memory/clip/embedder_test.go
Normal file
154
internal/memory/clip/embedder_test.go
Normal file
@ -0,0 +1,154 @@
|
|||||||
|
//go:build onnxruntime
|
||||||
|
|
||||||
|
package clip
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"fmt"
|
||||||
|
"image"
|
||||||
|
"image/color"
|
||||||
|
"image/png"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSmokeLoadAndEncode(t *testing.T) {
|
||||||
|
modelDir := os.Getenv("CLIP_MODEL_DIR")
|
||||||
|
if modelDir == "" {
|
||||||
|
modelDir = "/home/newqqagent/models/clip-vit-b32"
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(modelDir + "/text.onnx"); err != nil {
|
||||||
|
t.Skipf("模型目录不存在: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
emb, err := New(modelDir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("New: %v", err)
|
||||||
|
}
|
||||||
|
defer emb.Close()
|
||||||
|
|
||||||
|
if !emb.Loaded() {
|
||||||
|
t.Fatal("loaded should be true")
|
||||||
|
}
|
||||||
|
if emb.Dim() != 512 {
|
||||||
|
t.Fatalf("dim = %d, want 512", emb.Dim())
|
||||||
|
}
|
||||||
|
if emb.Fingerprint() == "" {
|
||||||
|
t.Fatal("fingerprint should not be empty")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 文本编码
|
||||||
|
textVec, err := emb.VectorizeDense("a photo of a cat")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("VectorizeDense: %v", err)
|
||||||
|
}
|
||||||
|
if len(textVec) != 512 {
|
||||||
|
t.Fatalf("text vec len = %d, want 512", len(textVec))
|
||||||
|
}
|
||||||
|
fmt.Printf("text vec[:5] = %v\n", textVec[:5])
|
||||||
|
|
||||||
|
// 同义文本应比远义文本更相似
|
||||||
|
textVec2, _ := emb.VectorizeDense("a photograph of a dog")
|
||||||
|
textVec3, _ := emb.VectorizeDense("quantum physics equations")
|
||||||
|
|
||||||
|
sim12 := cosineSim(textVec, textVec2)
|
||||||
|
sim13 := cosineSim(textVec, textVec3)
|
||||||
|
fmt.Printf("cat vs dog = %.4f, cat vs physics = %.4f\n", sim12, sim13)
|
||||||
|
if sim12 <= sim13 {
|
||||||
|
t.Errorf("cat-dog sim (%.4f) should be > cat-physics sim (%.4f)", sim12, sim13)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 通过 Vectorizer 接口(稀疏 map)
|
||||||
|
sparseVec := emb.Vectorize("hello world")
|
||||||
|
if len(sparseVec) == 0 {
|
||||||
|
t.Error("sparse Vectorize should return non-empty")
|
||||||
|
}
|
||||||
|
fmt.Printf("sparse len = %d\n", len(sparseVec))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCrossModalAlignment 验证图文在同一向量空间可比:
|
||||||
|
// 红底图的向量应与 "a red image" 更相似,而非 "a blue image"。
|
||||||
|
func TestCrossModalAlignment(t *testing.T) {
|
||||||
|
modelDir := os.Getenv("CLIP_MODEL_DIR")
|
||||||
|
if modelDir == "" {
|
||||||
|
modelDir = "/home/newqqagent/models/clip-vit-b32"
|
||||||
|
}
|
||||||
|
emb, err := New(modelDir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("New: %v", err)
|
||||||
|
}
|
||||||
|
defer emb.Close()
|
||||||
|
|
||||||
|
// 生成 224x224 纯红底 PNG
|
||||||
|
img := image.NewRGBA(image.Rect(0, 0, 224, 224))
|
||||||
|
red := color.RGBA{220, 40, 40, 255}
|
||||||
|
for y := 0; y < 224; y++ {
|
||||||
|
for x := 0; x < 224; x++ {
|
||||||
|
img.Set(x, y, red)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var buf bytes.Buffer
|
||||||
|
if err := png.Encode(&buf, img); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
imgVec, err := emb.EmbedImageDense(buf.Bytes(), "image/png")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("EmbedImageDense: %v", err)
|
||||||
|
}
|
||||||
|
if len(imgVec) != 512 {
|
||||||
|
t.Fatalf("img vec len = %d, want 512", len(imgVec))
|
||||||
|
}
|
||||||
|
|
||||||
|
redText, _ := emb.VectorizeDense("a red image")
|
||||||
|
blueText, _ := emb.VectorizeDense("a blue image")
|
||||||
|
redSim := cosineSim(imgVec, redText)
|
||||||
|
blueSim := cosineSim(imgVec, blueText)
|
||||||
|
fmt.Printf("red-image vs red-text = %.4f, vs blue-text = %.4f\n", redSim, blueSim)
|
||||||
|
if redSim <= blueSim {
|
||||||
|
t.Errorf("red image should align better with red text (%.4f) than blue (%.4f)", redSim, blueSim)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 同图应比异图更相似:存两张不同颜色,query 用红底图应召回红图
|
||||||
|
d1 := imgVec
|
||||||
|
blueImg := image.NewRGBA(image.Rect(0, 0, 224, 224))
|
||||||
|
blue := color.RGBA{40, 40, 220, 255}
|
||||||
|
for y := 0; y < 224; y++ {
|
||||||
|
for x := 0; x < 224; x++ {
|
||||||
|
blueImg.Set(x, y, blue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var buf2 bytes.Buffer
|
||||||
|
png.Encode(&buf2, blueImg)
|
||||||
|
d2, _ := emb.EmbedImageDense(buf2.Bytes(), "image/png")
|
||||||
|
if cosineSim(d1, d2) >= 0.99 {
|
||||||
|
t.Errorf("red and blue images should differ (got sim %.4f)", cosineSim(d1, d2))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cosineSim(a, b []float64) float64 {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
var dot, na, nb float64
|
||||||
|
for i := range a {
|
||||||
|
dot += a[i] * b[i]
|
||||||
|
na += a[i] * a[i]
|
||||||
|
nb += b[i] * b[i]
|
||||||
|
}
|
||||||
|
if na == 0 || nb == 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return dot / (sqrt(na) * sqrt(nb))
|
||||||
|
}
|
||||||
|
|
||||||
|
func sqrt(x float64) float64 {
|
||||||
|
if x <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
z := x
|
||||||
|
for i := 0; i < 50; i++ {
|
||||||
|
z = (z + x/z) / 2
|
||||||
|
}
|
||||||
|
return z
|
||||||
|
}
|
||||||
@ -24,9 +24,9 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"math"
|
"math"
|
||||||
"sort"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@ -646,6 +646,34 @@ func (s *Store) SetVec(digest string, vec []float64, model string) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// StaleVecDigests 返回所有需要重新嵌入的图片 digest:
|
||||||
|
// vec_model 不等于 currentModel(模型切换)或 vec_model 为空(从未嵌入)。
|
||||||
|
// 调用方使用返回的 digest 列表调用 Get/EmbedImage/SetVec 完成重算。
|
||||||
|
func (s *Store) StaleVecDigests(currentModel string) ([]string, error) {
|
||||||
|
s.mu.RLock()
|
||||||
|
defer s.mu.RUnlock()
|
||||||
|
|
||||||
|
rows, err := s.db.Query(`
|
||||||
|
SELECT digest FROM media
|
||||||
|
WHERE kind = 'image'
|
||||||
|
AND COALESCE(description,'') != ''
|
||||||
|
AND (COALESCE(vec_model,'') = '' OR vec_model != ?)
|
||||||
|
ORDER BY last_seen`, currentModel)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
var digests []string
|
||||||
|
for rows.Next() {
|
||||||
|
var d string
|
||||||
|
if err := rows.Scan(&d); err == nil {
|
||||||
|
digests = append(digests, d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return digests, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
// QueryMedia 用查询向量对所有已嵌入媒体做余弦相似度检索,返回 topK 个最相似的 Item。
|
// QueryMedia 用查询向量对所有已嵌入媒体做余弦相似度检索,返回 topK 个最相似的 Item。
|
||||||
//
|
//
|
||||||
// 这是跨模态检索的关键:查询可以是图片也可以是文本(经文本向量化后调用此方法),
|
// 这是跨模态检索的关键:查询可以是图片也可以是文本(经文本向量化后调用此方法),
|
||||||
|
|||||||
@ -133,3 +133,43 @@ func TestSetVec_PersistsCorrectly(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestStaleVecDigests(t *testing.T) {
|
||||||
|
s := newTestStore(t, 0)
|
||||||
|
defer s.Close()
|
||||||
|
|
||||||
|
// 有描述且 vec_model 匹配 → 非 stale
|
||||||
|
d1, _ := s.Put([]byte("img1"), Item{MIME: "image/png", Description: "图一"})
|
||||||
|
s.SetVec(d1, []float64{0.1}, "clip-vit-b32")
|
||||||
|
|
||||||
|
// 有描述但 vec_model 旧 → stale
|
||||||
|
d2, _ := s.Put([]byte("img2"), Item{MIME: "image/png", Description: "图二"})
|
||||||
|
s.SetVec(d2, []float64{0.2}, "clip-vit-b14")
|
||||||
|
|
||||||
|
// 有描述但从未嵌入(vec_model 空)→ stale
|
||||||
|
d3, _ := s.Put([]byte("img3"), Item{MIME: "image/png", Description: "图三"})
|
||||||
|
|
||||||
|
// 无描述 → 不参与(描述流程外)
|
||||||
|
s.Put([]byte("img4"), Item{MIME: "image/png"})
|
||||||
|
|
||||||
|
// 音频不属于图片 → 不算 stale
|
||||||
|
s.Put([]byte("aud1"), Item{MIME: "audio/wav", Description: "语音"})
|
||||||
|
|
||||||
|
stale, err := s.StaleVecDigests("clip-vit-b32")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(stale) != 2 {
|
||||||
|
t.Fatalf("expected 2 stale digests (d2 旧模型 + d3 未嵌入), got %d: %v", len(stale), stale)
|
||||||
|
}
|
||||||
|
got := map[string]bool{}
|
||||||
|
for _, d := range stale {
|
||||||
|
got[d] = true
|
||||||
|
}
|
||||||
|
if !got[d2] || !got[d3] {
|
||||||
|
t.Errorf("expected d2 and d3 stale, got %v", stale)
|
||||||
|
}
|
||||||
|
if got[d1] {
|
||||||
|
t.Errorf("d1 (匹配模型) 不应 stale")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@ -17,6 +17,18 @@ type Vectorizer interface {
|
|||||||
EmbedImage(img []byte, mime string) (Vector, error)
|
EmbedImage(img []byte, mime string) (Vector, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// MultimodalEmbedder 扩展 Vectorizer,提供直接返回 dense 向量的方法
|
||||||
|
// 与模型生命周期管理。CLIP 等视觉向量化器实现此接口;未启用时用空 stub。
|
||||||
|
type MultimodalEmbedder interface {
|
||||||
|
Vectorizer
|
||||||
|
VectorizeDense(text string) ([]float64, error)
|
||||||
|
EmbedImageDense(img []byte, mime string) ([]float64, error)
|
||||||
|
Fingerprint() string
|
||||||
|
Dim() int
|
||||||
|
Loaded() bool
|
||||||
|
Close()
|
||||||
|
}
|
||||||
|
|
||||||
// ErrNotSupported 表示 Vectorizer 不支持图像嵌入,调用方按文本描述降级。
|
// ErrNotSupported 表示 Vectorizer 不支持图像嵌入,调用方按文本描述降级。
|
||||||
var ErrNotSupported = fmt.Errorf("vectorizer does not support image embedding")
|
var ErrNotSupported = fmt.Errorf("vectorizer does not support image embedding")
|
||||||
|
|
||||||
@ -256,7 +268,7 @@ func CosineSimilarity(a, b Vector) float64 {
|
|||||||
|
|
||||||
// InvertedIndex 倒排索引,加速向量搜索
|
// InvertedIndex 倒排索引,加速向量搜索
|
||||||
type InvertedIndex struct {
|
type InvertedIndex struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
postings map[string]map[string]float64 // feature → {docID: weight}
|
postings map[string]map[string]float64 // feature → {docID: weight}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user