diff --git a/internal/memory/media/media.go b/internal/memory/media/media.go index dc7338d..9b47347 100644 --- a/internal/memory/media/media.go +++ b/internal/memory/media/media.go @@ -23,6 +23,8 @@ import ( "encoding/json" "fmt" "io" + "math" + "sort" "os" "path/filepath" "strings" @@ -91,6 +93,12 @@ type Item struct { // FirstSeen/LastSeen 是首末次入库时间。 FirstSeen time.Time `json:"first_seen"` LastSeen time.Time `json:"last_seen"` + // --- 多模态嵌入(v1.2.0)--- + // Vec 是视觉嵌入向量的序列化(JSON []float64),nil 表示未嵌入。 + Vec []float64 `json:"vec,omitempty"` + // VecModel 是产生 Vec 的模型标识(如 "clip-vit-b32"), + // 用于模型切换后判断是否需要重算。 + VecModel string `json:"vec_model,omitempty"` } // Store 管理媒体的元数据(SQLite)与内容(磁盘 CAS 目录)。 @@ -166,6 +174,14 @@ func (s *Store) initSchema() error { return fmt.Errorf("media: schema %q: %w", truncate(q, 60), err) } } + // v1.2.0 迁移:给 media 表加 vec(视觉嵌入向量 JSON)和 vec_model(模型标识)。 + migrations := []string{ + `ALTER TABLE media ADD COLUMN vec TEXT`, + `ALTER TABLE media ADD COLUMN vec_model TEXT`, + } + for _, q := range migrations { + _, _ = s.db.Exec(q) // 列已存在时返回 "duplicate column name",可忽略 + } return nil } @@ -609,6 +625,117 @@ func (s *Store) Close() error { return s.db.Close() } +// ---- 多模态嵌入(v1.2.0) ---- + +// SetVec 给一条已入库的媒体设置视觉嵌入向量。 +// +// 设计选择:vec 是 TEXT(JSON 序列化的 []float64)而非 BLOB, +// 因为 Go 的 json.Marshal/Unmarshal 对 []float64 是自然的, +// 而 SQLite 的 BLOB 是 []byte,序列化多一层反而复杂。 +// 量级:一条 vec 最多 1536 维 × ~15 字节 ≈ 23KB,TEXT 合适。 +func (s *Store) SetVec(digest string, vec []float64, model string) error { + vecJSON, err := json.Marshal(vec) + if err != nil { + return fmt.Errorf("media: marshal vec: %w", err) + } + s.mu.Lock() + defer s.mu.Unlock() + _, err = s.db.Exec(`UPDATE media SET vec = ?, vec_model = ? WHERE digest = ?`, + string(vecJSON), model, digest) + return err +} + +// QueryMedia 用查询向量对所有已嵌入媒体做余弦相似度检索,返回 topK 个最相似的 Item。 +// +// 这是跨模态检索的关键:查询可以是图片也可以是文本(经文本向量化后调用此方法), +// 被查的媒体库里的每个 item 也有一个视觉向量。两者在同一空间比对, +// 谁的相似度更高就召回谁——不再区分「这是一张图的查询」还是「这是一段文字的查询」, +// 由向量空间的相似度自动判断。 +func (s *Store) QueryMedia(queryVec []float64, model string, topK int) ([]*Item, error) { + if topK <= 0 { + topK = 20 + } + if len(queryVec) == 0 { + return nil, nil + } + s.mu.RLock() + defer s.mu.RUnlock() + + rows, err := s.db.Query(`SELECT digest, kind, mime, size, width, height, + origin_path, tool, description, described_by, ref_count, first_seen, last_seen, + vec, vec_model + FROM media WHERE vec IS NOT NULL AND vec != ''`) + if err != nil { + return nil, err + } + defer rows.Close() + + type scored struct { + item *Item + score float64 + } + var candidates []scored + for rows.Next() { + var it Item + var kind string + var origin, tool, desc, by, vecJSON, vecModel sql.NullString + if err := rows.Scan(&it.Digest, &kind, &it.MIME, &it.Size, &it.Width, &it.Height, + &origin, &tool, &desc, &by, &it.RefCount, &it.FirstSeen, &it.LastSeen, + &vecJSON, &vecModel); err != nil { + continue + } + it.Kind = Kind(kind) + it.OriginPath = origin.String + it.Tool = tool.String + it.Description = desc.String + it.DescribedBy = by.String + if !vecJSON.Valid || vecJSON.String == "" { + continue + } + var itemVec []float64 + if err := json.Unmarshal([]byte(vecJSON.String), &itemVec); err != nil || len(itemVec) == 0 { + continue + } + if len(itemVec) != len(queryVec) { + continue // 维度不一致,跳过 + } + score := cosineSimilaritySlice(queryVec, itemVec) + if score > 0.05 { + candidates = append(candidates, scored{&it, score}) + } + } + if err := rows.Err(); err != nil { + return nil, err + } + + // 按分数降序排序 + sort.Slice(candidates, func(i, j int) bool { + return candidates[i].score > candidates[j].score + }) + if len(candidates) > topK { + candidates = candidates[:topK] + } + out := make([]*Item, len(candidates)) + for i, c := range candidates { + out[i] = c.item + } + return out, nil +} + +// cosineSimilaritySlice 计算两个 []float64 向量的余弦相似度。 +func cosineSimilaritySlice(a, b []float64) float64 { + var dot, normA, normB float64 + for i := range a { + dot += a[i] * b[i] + normA += a[i] * a[i] + normB += b[i] * b[i] + } + if normA == 0 || normB == 0 { + return 0 + } + return dot / (math.Sqrt(normA) * math.Sqrt(normB)) +} + // ---- 扫描辅助 ---- type rowScanner interface { @@ -628,9 +755,10 @@ func (s *Store) scanRows(r rowScanner) (*Item, error) { return scanItem(r) } func scanItem(r rowScanner) (*Item, error) { var it Item var kind string - var origin, tool, desc, by sql.NullString + var origin, tool, desc, by, vecJSON, vecModel sql.NullString if err := r.Scan(&it.Digest, &kind, &it.MIME, &it.Size, &it.Width, &it.Height, - &origin, &tool, &desc, &by, &it.RefCount, &it.FirstSeen, &it.LastSeen); err != nil { + &origin, &tool, &desc, &by, &it.RefCount, &it.FirstSeen, &it.LastSeen, + &vecJSON, &vecModel); err != nil { return nil, err } it.Kind = Kind(kind) @@ -638,6 +766,13 @@ func scanItem(r rowScanner) (*Item, error) { it.Tool = tool.String it.Description = desc.String it.DescribedBy = by.String + if vecJSON.Valid && vecJSON.String != "" { + var v []float64 + if err := json.Unmarshal([]byte(vecJSON.String), &v); err == nil { + it.Vec = v + } + } + it.VecModel = vecModel.String return &it, nil }