这套方案做完之后,你会得到一个完全跑在本机的图片检索小系统,能回答三类问题:
- 文搜图:输入「会议室白板上的架构图」,返回相册里最像的几张图。
- 以图搜图:丢一张截图进去,找出同一批文件里视觉上相近的图。
- 图文混检:同时给一段文字和一张示例图,两路结果融合后统一排序。
图片不出本机,不调任何云端接口,普通笔记本的 CPU 也能跑起来,加一块 GPU 会明显更快。检索范围覆盖几千到几万张图片的常见个人/团队素材库。
先理清一件事:跨模态检索为什么需要两个塔
这是整篇教程里最容易踩的认知坑,先说清楚。
跨模态检索的本质要求是:查询向量和候选向量必须落在同一个向量空间里,才可以用AI 词典:余弦相似度">余弦相似度比较。一段文字想要直接和一张图片比大小,就得有一个「在图文对上训练过」的编码器,把图和文映射到同一个空间。
EmbeddingGemma 系列是轻量级的嵌入模型,主打本地可跑、体积小,主要面向文本嵌入。所以这套系统这样分工:
- 视觉塔:负责图片向量,以及和图片同空间的文本向量(文搜图、以图搜图走这条路)。
- EmbeddingGemma:负责把每张图的描述文本(文件名、目录名、标签、caption)变成向量,处理「同义改写」「中文长句」这类语义检索(文搜描述走这条路)。
- 融合层:用 RRF(倒数排名融合)把两路排序合并,得到混检结果。
模型是否自带图文对齐能力、模型 ID 具体叫什么、有没有对应的 query/document prompt 前缀,一律以官方模型卡当前版本为准。如果官方后续放出的版本本身就带图文对齐,把下面的视觉塔换成官方权重即可,其余流程完全不用改。
前置条件清单
- Python 3.10 及以上(具体支持的版本以官方文档为准)
- 一个图片文件夹,格式为 jpg / png / webp / bmp 之类常见格式
- 磁盘预留出模型权重的空间(嵌入模型通常不大,视觉塔略大,具体看模型卡)
- 依赖包:
torch、transformers、sentence-transformers、faiss-cpu、Pillow、tqdm - 可选:GPU(装对应 CUDA 版本的 torch 即可,不装也能跑)
目录结构约定如下,后面所有路径都按这个来:
```text
image-search/
├── images/ # 你的图片素材
├── models/ # 模型权重本地缓存
├── index/ # 向量库 + id 映射
├── captions.json # 可选:每张图的描述文本
├── config.py
├── build_index.py
└── search.py
```
第一步:建环境、装依赖
```bash
python -m venv .venv
source .venv/bin/activate # Windows: .venv\Scripts\activate
python -m pip install -U pip
python -m pip install torch transformers sentence-transformers faiss-cpu Pillow tqdm
```
GPU 用户把 torch 换成对应 CUDA 版本(安装命令以 PyTorch 官方页面为准)。
第二步:把模型拉到本地
先设一个本地缓存目录,避免默认缓存散落在用户目录里,也方便之后离线跑。
```bash
export HF_HOME="$PWD/models"
python -m pip install -U "huggingface_hub[cli]"
```
这两个模型先拉到本地目录(把占位符换成官方模型卡上的实际 ID):
```bash
huggingface-cli download <文本塔模型ID> --local-dir "$HF_HOME/text-tower"
huggingface-cli download <视觉塔模型ID> --local-dir "$HF_HOME/vision-tower"
```
新版 huggingface_hub 里命令改叫 hf download,两个名字任选其一,以官方文档当前版本为准。国内网络环境下可以设置 HF_ENDPOINT 指向镜像站;拉完模型之后想彻底离线跑,设置 HF_HUB_OFFLINE=1。
第三步:配置文件
```python
config.py
from pathlib import Path
ROOT = Path(__file__).resolve().parent
IMAGE_DIR = ROOT / "images"
CACHE_DIR = ROOT / "models"
INDEX_DIR = ROOT / "index"
文本塔:EmbeddingGemma,ID 以官方模型卡当前版本为准
TEXT_MODEL_ID = "REPLACE_WITH_OFFICIAL_TEXT_MODEL_ID"
视觉塔:任意图文对齐的开源编码器(CLIP / SigLIP 系列等),同样以官方模型卡为准
VISION_MODEL_ID = "REPLACE_WITH_OFFICIAL_VISION_MODEL_ID"
DEVICE = "cuda" # 没有 GPU 就改成 "cpu"
IMAGE_BATCH = 16
TEXT_BATCH = 32
```
第四步:图片批量向量化
这一步的要点是流式分批:不要一次把所有图片读进内存,几万张图会直接把内存吃满。
```python
build_index.py
import json
from pathlib import Path
import faiss
import numpy as np
import torch
from PIL import Image, ImageOps
from tqdm import tqdm
from transformers import CLIPModel, CLIPProcessor
from sentence_transformers import SentenceTransformer
import config as C
IMG_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp", ".tif", ".tiff"}
def list_images(root: Path):
root = Path(root)
return sorted(p for p in root.rglob("*") if p.suffix.lower() in IMG_EXTS)
def load_rgb(path: Path) -> Image.Image:
"""统一转成 RGB,顺手处理手机照片的 EXIF 旋转和透明底。"""
im = Image.open(path)
im = ImageOps.exif_transpose(im)
if im.mode in ("RGBA", "LA", "P"):
im = im.convert("RGBA")
bg = Image.new("RGB", im.size, (255, 255, 255))
bg.paste(im, mask=im.split()[-1])
return bg
return im.convert("RGB")
def describe(path: Path, captions: dict) -> str:
"""每张图对应的描述文本:优先用 captions.json,否则退回目录名 + 文件名。"""
key = str(path.relative_to(C.IMAGE_DIR))
if key in captions:
return captions[key]
parts = [path.parent.name,
path.stem.replace("_", " ").replace("-", " ")]
return " ".join(p for p in parts if p)
def save_index(vecs: np.ndarray, out_path: Path) -> None:
"""内积 + L2 归一化 = 余弦相似度。"""
vecs = np.ascontiguousarray(vecs, dtype="float32")
faiss.normalize_L2(vecs)
index = faiss.IndexFlatIP(vecs.shape[1])
index.add(vecs)
faiss.write_index(index, str(out_path))
```
然后是主流程:
```python
def main():
C.INDEX_DIR.mkdir(parents=True, exist_ok=True)
paths = list_images(C.IMAGE_DIR)
print(f"共找到 {len(paths)} 张图片")
if not paths:
return
cap_file = C.ROOT / "captions.json"
captions = json.loads(cap_file.read_text(encoding="utf-8")) if cap_file.exists() else {}
---- 视觉塔:图像向量 ----
vproc = CLIPProcessor.from_pretrained(C.VISION_MODEL_ID, cache_dir=str(C.CACHE_DIR))
vmodel = CLIPModel.from_pretrained(C.VISION_MODEL_ID, cache_dir=str(C.CACHE_DIR))
vmodel = vmodel.to(C.DEVICE).eval()
img_vecs = []
with torch.no_grad():
for i in tqdm(range(0, len(paths), C.IMAGE_BATCH), desc="图片向量化"):
chunk = paths[i:i + C.IMAGE_BATCH]
imgs = [load_rgb(p) for p in chunk]
inputs = vproc(images=imgs, return_tensors="pt").to(C.DEVICE)
feats = vmodel.get_image_features(**inputs)
feats = torch.nn.functional.normalize(feats.float(), dim=-1)
img_vecs.append(feats.cpu().numpy())
img_vecs = np.concatenate(img_vecs, axis=0)
---- 文本塔:描述文本向量 ----
docs = [describe(p, captions) for p in paths]
tmodel = SentenceTransformer(C.TEXT_MODEL_ID,
cache_folder=str(C.CACHE_DIR),
device=C.DEVICE)
doc_vecs = tmodel.encode(docs,
batch_size=C.TEXT_BATCH,
normalize_embeddings=True,
convert_to_numpy=True,
show_progress_bar=True)
save_index(img_vecs, C.INDEX_DIR / "image.faiss")
save_index(doc_vecs, C.INDEX_DIR / "doc.faiss")
with (C.INDEX_DIR / "meta.jsonl").open("w", encoding="utf-8") as f:
for i, (p, d) in enumerate(zip(paths, docs)):
f.write(json.dumps({
"idx": i,
"path": str(p.relative_to(C.ROOT)), # 存相对路径,换机器也能用
"doc": d,
}, ensure_ascii=False) + "\n")
print("索引构建完成")
if __name__ == "__main__":
main()
```
注意两点:
1. 视觉塔这里用了 CLIPModel / CLIPProcessor。如果选的是 SigLIP 等其他家族,类名不同,以对应官方文档为准。
2. 如果文本塔的模型卡给出了 query / document 的 prompt 前缀,encode 时用 prompt_name= 分别传入,查询侧和文档侧不要用同一个,否则召回会明显变差。
跑起来:
```bash
python build_index.py
```
第五步:检索脚本
```python
search.py
import argparse
import json
from pathlib import Path
import faiss
import numpy as np
import torch
from PIL import Image
from transformers import CLIPModel, CLIPProcessor
from sentence_transformers import SentenceTransformer
import config as C
from build_index import load_rgb
def load_meta():
rows = [json.loads(l) for l in
(C.INDEX_DIR / "meta.jsonl").read_text(encoding="utf-8").splitlines() if l.strip()]
return rows
def rrf(rank_lists, weights=None, k=60):
"""倒数排名融合:只关心名次,不关心两路分数是否可比。"""
scores = {}
for i, ranks in enumerate(rank_lists):
w = 1.0 if weights is None else weights[i]
for r, doc in enumerate(ranks):
scores[doc] = scores.get(doc, 0.0) + w / (k + r + 1)
return [d for d, _ in sorted(scores.items(), key=lambda kv: -kv[1])]
def top_ids(index, query_vec, topk):
q = np.ascontiguousarray(query_vec.astype("float32"))[None, :]
faiss.normalize_L2(q)
_, ids = index.search(q, topk)
return [int(i) for i in ids[0] if i >= 0]
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--text", default=None, help="文字查询")
ap.add_argument("--image", default=None, help="示例图片路径(以图搜图)")
ap.add_argument("--topk", type=int, default=8)
args = ap.parse_args()
meta = load_meta()
img_index = faiss.read_index(str(C.INDEX_DIR / "image.faiss"))
doc_index = faiss.read_index(str(C.INDEX_DIR / "doc.faiss"))
vproc = CLIPProcessor.from_pretrained(C.VISION_MODEL_ID, cache_dir=str(C.CACHE_DIR))
vmodel = CLIPModel.from_pretrained(C.VISION_MODEL_ID,
cache_dir=str(C.CACHE_DIR)).to(C.DEVICE).eval()
rank_lists, labels = [], []
A 路:文字 -> 视觉塔文本编码 -> 图像向量库(真正的文搜图)
if args.text:
with torch.no_grad():
inputs = vproc(text=[args.text], return_tensors="pt",
padding=True, truncation=True).to(C.DEVICE)
q = vmodel.get_text_features(**inputs).float()
q = torch.nn.functional.normalize(q, dim=-1).cpu().numpy()[0]
rank_lists.append(top_ids(img_index, q, args.topk))
labels.append("文搜图")
B 路:文字 -> EmbeddingGemma -> 描述向量库(语义检索)
tmodel = SentenceTransformer(C.TEXT_MODEL_ID,
cache_folder=str(C.CACHE_DIR),
device=C.DEVICE)
qd = tmodel.encode([args.text], normalize_embeddings=True,
convert_to_numpy=True)[0]
rank_lists.append(top_ids(doc_index, qd, args.topk))
labels.append("文搜描述")
C 路:图片 -> 视觉塔图像编码 -> 图像向量库(以图搜图)
if args.image:
with torch.no_grad():
inputs = vproc(images=[load_rgb(Path(args.image))],
return_tensors="pt").to(C.DEVICE)
q = vmodel.get_image_features(**inputs).float()
q = torch.nn.functional.normalize(q, dim=-1).cpu().numpy()[0]
rank_lists.append(top_ids(img_index, q, args.topk))
labels.append("以图搜图")
if not rank_lists:
print("至少给一个 --text 或 --image")
return
fused = rrf(rank_lists)[:args.topk]
print(f"融合来源:{', '.join(labels)}")
for rank, idx in enumerate(fused, 1):
print(f"{rank:2d}. {meta[idx]['path']} [{(meta[idx]['doc'] or '')[:40]}]")
if __name__ == "__main__":
main()
```
三种用法:
```bash
文搜图(同时走视觉塔文本塔和描述库,RRF 融合)
python search.py --text "会议室白板上的架构图"
以图搜图
python search.py --image ./query.jpg --topk 12
图文混检
python search.py --text "发票扫描件" --image ./sample.jpg --topk 10
```
常见坑与排错
两路向量混进同一个索引。 视觉塔的图像向量和 EmbeddingGemma 的文本向量维度可能碰巧相同,但语义空间完全不同,放进同一个 IndexFlatIP 里比较毫无意义。必须分成两个索引,靠 RRF 在排序层面融合。
忘了归一化。 用 IndexFlatIP 算内积,只有在向量都是单位长度时才等价于余弦相似度。图片侧和查询侧都要做一次 L2 归一化,漏掉任何一边,排名都会偏。
文字查询走了错的塔。 想「文字找图」,查询文本必须用视觉塔自己的文本编码器编码;用 EmbeddingGemma 编码出来的向量和图像向量不在一个空间。EmbeddingGemma 负责的是检索描述文本,再把命中的图片捞出来。
透明底图片变成黑块。 PNG 带 alpha 通道时,直接 convert("RGB") 会把透明区域填黑,和实际观感差距很大。上面 load_rgb 里贴到白底上就是为了这个。
手机照片方向不对。 EXIF 里存了旋转信息,不调 ImageOps.exif_transpose 的话,竖拍照片会横着进模型,检索质量下降。
描述文本全空。 如果 captions.json 不存在,而文件名又是 IMG_0001.jpg 这种,B 路基本等于没用。建议至少给文件夹起有意义的名字,或者补一份 caption。批量生成 caption 可以用 BLIP 一类模型或人工标注,具体选型按需求定。
路径存成了绝对路径。 索引里存绝对路径,换台机器或者改了目录名就全失效。上面统一存相对路径,查询时再拼回根目录。
中文查询效果差。 如果视觉塔只在英文图文对上训练过,中文查询会比较吃力。可以考虑给每张图配中文描述走 B 路,或者换一个支持中文的视觉塔,具体以模型卡说明为准。
下载卡住或反复重下。 统一用 HF_HOME 指向本地目录,拉完之后设 HF_HUB_OFFLINE=1 强制离线,避免每次启动都去连网络。
图片太多内存爆掉。 检查是不是把全部图片对象一次性读进了列表。批处理循环里每批用完就该释放,别在循环外保留引用。
FAISS 报类型错误。 faiss 只认 float32,从 PyTorch 转过来时记得 .astype("float32"),float16 会直接报错。
下一步可以做什么
- 加一个重排序环节:先粗排取前 50,再用更强的模型精排取前 10,召回质量和排序质量都会改善。
- 做增量索引:
IndexFlatIP支持直接add新向量,把文件哈希存下来,重复图片和已索引图片直接跳过。 - 接 OCR:截图、发票、PPT 导出图里大量信息是文字,加一路 OCR 文本进 EmbeddingGemma 索引,检索命中率会明显提升。
- 包一层界面:用 Gradio 或 FastAPI 起个小服务,上传图片就能查,团队内共享更方便。
- 上量化与 ONNX:想塞进更小的设备或提速,导出为 ONNX 或做整数量化,具体支持情况以官方文档为准。
- 换更强的塔:数据量涨上去之后,视觉塔和文本塔都可以换成更大的版本;只要索引重建流程是脚本化的,替换成本很低。
整套东西的核心其实就一句话:让该在一个空间里的向量待在一起,不该在一起的用排序融合去合。把握住这条,模型怎么换、库怎么升级,你都不会迷路。
