跳到主内容
快讯直播
AI智模界
教程

用 EmbeddingGemma 2 搭一个图文跨模态的本地检索

这套方案做完之后,你会得到一个完全跑在本机的图片检索小系统,能回答三类问题:

  • 文搜图:输入「会议室白板上的架构图」,返回相册里最像的几张图。
  • 以图搜图:丢一张截图进去,找出同一批文件里视觉上相近的图。
  • 图文混检:同时给一段文字和一张示例图,两路结果融合后统一排序。

图片不出本机,不调任何云端接口,普通笔记本的 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 或做整数量化,具体支持情况以官方文档为准。
  • 换更强的塔:数据量涨上去之后,视觉塔和文本塔都可以换成更大的版本;只要索引重建流程是脚本化的,替换成本很低。

整套东西的核心其实就一句话:让该在一个空间里的向量待在一起,不该在一起的用排序融合去合。把握住这条,模型怎么换、库怎么升级,你都不会迷路。

AI 生成本文由 AI 基于公开信息自动生成,仅供参考。