跨模态检索的工程挑战:文本搜图、图搜文本和图搜图的统一架构

跨模态检索的工程挑战:文本搜图、图搜文本和图搜图的统一架构

一、深度引言与场景痛点

大家好,我是赵咕咕。

电商设计团队的一个日常需求是这样的:运营给一张竞品图,说"帮我找一下我们库里跟这个风格最接近的 banner 素材"。设计师可能再补一句"色调尽量类似上周那个蓝色渐变方案的"。

这个需求看起来简单,拆开来看包含三种不同的检索模式:

  1. 文本搜图:用文字描述图片内容,找到匹配的图片("蓝色渐变 banner")。
  2. 图搜文本:给一张图,找到描述这张图的文字(找到对应的设计说明、prompt 或标签)。
  3. 图搜图:给一张图,找到视觉上相似的图(竞品图 → 库内相似素材)。

三种模式对应三个不同的"检索方向"。传统做法是三套独立系统,但本质上它们应该共享同一个理解模型——因为无论是文字描述一张图,还是判断两张图是否相似,底层都是对视觉语义的理解。

这篇文章,我聊聊如何在一个统一架构里支持这三种跨模态检索,核心依靠 CLIP 模型将图文映射到同一向量空间。

二、底层机制与原理深度剖析

2.1 CLIP 的核心思想

CLIP(Contrastive Language-Image Pre-training)的核心思路很优雅:用 4 亿对图文配对数据训练,让模型学会把匹配的图文对拉近,不匹配的推远。

训练完成后,CLIP 有两个输出分支:

  • 文本编码器:任意文本 → 固定维度向量(如 512 维)
  • 图像编码器:任意图片 → 同一维度向量

因为它们在同一个向量空间里,所以:

  • 文本搜图 = 在图片向量库中找跟查询文本向量最近的图片
  • 图搜文本 = 在文本向量库中找跟查询图片向量最近的文本
  • 图搜图 = 在图片向量库中找跟查询图片向量最近的图片

这就是统一架构的数学基础——不需要三套模型,一个 CLIP 模型就够了。

2.2 统一架构设计

架构的四个关键层:

  1. 共享编码层:文本和图像经过各自的编码器,进入同一个 L2 归一化的向量空间。这是统一架构的基础。
  2. 多模态索引:图片特征和文本特征存在同一个向量数据库中(使用相同的距离度量),只是 collection 分开。
  3. 查询路由层:根据输入类型,路由到不同的检索路径。text→搜图,image→搜图,image→搜文本。
  4. 后处理层:元数据过滤、去重、可解释性。特别是"为什么这张图跟查询匹配"的解释——对于创意工作者来说,知道"因为色调相似所以推荐了这张"比知道"相似度 0.87"有用得多。

2.3 混合查询的加权融合

最实用的场景是"模糊文本描述 + 一张参考图"的混合查询:

  • "帮我找跟这张图(上传竞品截图)风格类似的素材"
  • 或者"大概长这样(文字描述),色调参考这张(上传参考图)"

处理方式是:文本编码向量 × α + 图片编码向量 × (1-α) → 融合向量 → 搜索图片索引。

α 是一个可调节的权重。α=0.7 表示"更重视文字描述的语义",α=0.3 表示"更接近上传图片的风格"。业务上可以让用户在 UI 中拖动滑块调整 α。

三、生产级代码实现

import asyncio import base64 import hashlib import logging from dataclasses import dataclass, field from enum import Enum from typing import Any import numpy as np import torch from PIL import Image from qdrant_client import QdrantClient from qdrant_client.models import ( Distance, VectorParams, PointStruct, Filter, FieldCondition, MatchValue, ) # 通过 sentence-transformers 或 OpenCLIP 加载 CLIP from sentence_transformers import SentenceTransformer logger = logging.getLogger(__name__) # ── 数据模型 ─────────────────────────────────────────── class SearchMode(Enum): TEXT_TO_IMAGE = "text_to_image" IMAGE_TO_IMAGE = "image_to_image" IMAGE_TO_TEXT = "image_to_text" HYBRID = "hybrid" class Modality(Enum): IMAGE = "image" TEXT = "text" @dataclass class MediaAsset: """媒体资产模型。""" asset_id: str modality: Modality content: str # 文本内容 或 图片路径 tags: list[str] = field(default_factory=list) metadata: dict[str, Any] = field(default_factory=dict) embedding: np.ndarray | None = None @dataclass class CrossModalResult: """跨模态检索结果。""" asset: MediaAsset score: float modality: Modality source: str = "" # 来自哪个索引 # ── CLIP 编码器 ──────────────────────────────────────── class CLIPEncoder: """CLIP 模型封装,统一文本和图像编码。""" def __init__(self, model_name: str = "clip-ViT-B-32"): self._model_name = model_name self._model: SentenceTransformer | None = None self._dim = 512 async def initialize(self) -> None: if self._model is not None: return try: self._model = await asyncio.to_thread( SentenceTransformer, self._model_name ) # 获取模型输出维度 self._dim = self._model.get_sentence_embedding_dimension() logger.info("CLIP 模型 %s 加载完成, 维度=%d", self._model_name, self._dim) except Exception as e: logger.error("CLIP 模型加载失败: %s", e) raise async def encode_text(self, text: str) -> np.ndarray: """文本编码。""" assert self._model is not None, "模型未初始化" embedding = await asyncio.to_thread( self._model.encode, [text], normalize_embeddings=True, show_progress_bar=False ) return embedding[0] async def encode_image(self, image_path: str) -> np.ndarray: """图像编码。""" assert self._model is not None, "模型未初始化" try: img = Image.open(image_path).convert("RGB") embedding = await asyncio.to_thread( self._model.encode, [img], normalize_embeddings=True, show_progress_bar=False ) return embedding[0] except Exception as e: logger.error("图像编码失败 %s: %s", image_path, e) raise async def encode_image_base64(self, image_b64: str) -> np.ndarray: """从 base64 编码图像。""" import io img_data = base64.b64decode(image_b64) img = Image.open(io.BytesIO(img_data)).convert("RGB") embedding = await asyncio.to_thread( self._model.encode, [img], normalize_embeddings=True, show_progress_bar=False ) return embedding[0] @property def dimension(self) -> int: return self._dim # ── 统一的多模态检索引擎 ─────────────────────────────── class CrossModalSearchEngine: """跨模态检索引擎,支持 text→image, image→image, image→text 三种模式。""" IMAGE_COLLECTION = "media_images" TEXT_COLLECTION = "media_texts" def __init__(self, encoder: CLIPEncoder | None = None): self._encoder = encoder or CLIPEncoder() self._client = QdrantClient(path="./qdrant_multimodal") async def initialize(self) -> None: """初始化编码器和向量集合。""" await self._encoder.initialize() for coll in [self.IMAGE_COLLECTION, self.TEXT_COLLECTION]: if not self._client.collection_exists(coll): self._client.create_collection( collection_name=coll, vectors_config=VectorParams( size=self._encoder.dimension, distance=Distance.COSINE, ), ) # 创建标签索引 self._client.create_payload_index( collection_name=coll, field_name="tags", field_schema="keyword", ) logger.info("多模态检索引擎初始化完成") async def index_assets(self, assets: list[MediaAsset]) -> None: """批量索引媒体资产(图片和文本混合)。""" image_points = [] text_points = [] for asset in assets: try: if asset.modality == Modality.IMAGE: embedding = await self._encoder.encode_image(asset.content) point_id = hashlib.md5(asset.asset_id.encode()).hexdigest()[:16] image_points.append(PointStruct( id=point_id, vector=embedding.tolist(), payload={ "asset_id": asset.asset_id, "path": asset.content, "tags": asset.tags, **asset.metadata, }, )) else: embedding = await self._encoder.encode_text(asset.content) point_id = hashlib.md5(asset.asset_id.encode()).hexdigest()[:16] text_points.append(PointStruct( id=point_id, vector=embedding.tolist(), payload={ "asset_id": asset.asset_id, "content": asset.content, "tags": asset.tags, **asset.metadata, }, )) except Exception as e: logger.error("索引资产 %s 失败: %s", asset.asset_id, e) if image_points: self._client.upsert(collection_name=self.IMAGE_COLLECTION, points=image_points) logger.info("已索引 %d 张图片", len(image_points)) if text_points: self._client.upsert(collection_name=self.TEXT_COLLECTION, points=text_points) logger.info("已索引 %d 条文本", len(text_points)) async def search( self, query_text: str | None = None, query_image_path: str | None = None, query_image_b64: str | None = None, mode: SearchMode = SearchMode.TEXT_TO_IMAGE, hybrid_alpha: float = 0.5, filter_tags: list[str] | None = None, top_k: int = 10, ) -> list[CrossModalResult]: """统一的跨模态搜索入口。""" results = [] # 构建查询过滤器 query_filter = None if filter_tags: conditions = [ FieldCondition(key="tags", match=MatchValue(value=tag)) for tag in filter_tags ] query_filter = Filter(must=conditions) try: if mode == SearchMode.TEXT_TO_IMAGE and query_text: # 文本 → 图像 vec = await self._encoder.encode_text(query_text) hits = self._client.search( collection_name=self.IMAGE_COLLECTION, query_vector=vec.tolist(), query_filter=query_filter, limit=top_k, ) for hit in hits: p = hit.payload or {} results.append(CrossModalResult( asset=MediaAsset( asset_id=p.get("asset_id", ""), modality=Modality.IMAGE, content=p.get("path", ""), tags=p.get("tags", []), metadata=p, ), score=hit.score, modality=Modality.IMAGE, source="image_index", )) elif mode == SearchMode.IMAGE_TO_IMAGE and (query_image_path or query_image_b64): # 图像 → 图像 vec = ( await self._encoder.encode_image(query_image_path) if query_image_path else await self._encoder.encode_image_base64(query_image_b64) ) hits = self._client.search( collection_name=self.IMAGE_COLLECTION, query_vector=vec.tolist(), query_filter=query_filter, limit=top_k, ) for hit in hits: p = hit.payload or {} results.append(CrossModalResult( asset=MediaAsset( asset_id=p.get("asset_id", ""), modality=Modality.IMAGE, content=p.get("path", ""), tags=p.get("tags", []), metadata=p, ), score=hit.score, modality=Modality.IMAGE, source="image_index", )) elif mode == SearchMode.IMAGE_TO_TEXT and (query_image_path or query_image_b64): # 图像 → 文本 vec = ( await self._encoder.encode_image(query_image_path) if query_image_path else await self._encoder.encode_image_base64(query_image_b64) ) hits = self._client.search( collection_name=self.TEXT_COLLECTION, query_vector=vec.tolist(), query_filter=query_filter, limit=top_k, ) for hit in hits: p = hit.payload or {} results.append(CrossModalResult( asset=MediaAsset( asset_id=p.get("asset_id", ""), modality=Modality.TEXT, content=p.get("content", ""), tags=p.get("tags", []), metadata=p, ), score=hit.score, modality=Modality.TEXT, source="text_index", )) elif mode == SearchMode.HYBRID and query_text and (query_image_path or query_image_b64): # 混合查询:文本向量 × α + 图片向量 × (1-α) text_vec = await self._encoder.encode_text(query_text) img_vec = ( await self._encoder.encode_image(query_image_path) if query_image_path else await self._encoder.encode_image_base64(query_image_b64) ) # 加权融合(确保归一化) fused = text_vec * hybrid_alpha + img_vec * (1 - hybrid_alpha) fused = fused / np.linalg.norm(fused) hits = self._client.search( collection_name=self.IMAGE_COLLECTION, query_vector=fused.tolist(), query_filter=query_filter, limit=top_k, ) for hit in hits: p = hit.payload or {} results.append(CrossModalResult( asset=MediaAsset( asset_id=p.get("asset_id", ""), modality=Modality.IMAGE, content=p.get("path", ""), tags=p.get("tags", []), metadata=p, ), score=hit.score, modality=Modality.IMAGE, source=f"hybrid(α={hybrid_alpha})", )) except Exception as e: logger.exception("跨模态搜索失败: mode=%s", mode) return [] # 最大边缘相关度(MMR)去重:保证结果多样性 results = self._mmr_dedup(results, lambda_coef=0.5, final_k=min(top_k, len(results))) return results def _mmr_dedup( self, results: list[CrossModalResult], lambda_coef: float = 0.5, final_k: int = 5, ) -> list[CrossModalResult]: """MMR 去重:平衡相关性和多样性。""" if len(results) <= final_k: return results selected: list[CrossModalResult] = [] remaining = list(results) # 第一个选最高分的 remaining.sort(key=lambda x: x.score, reverse=True) selected.append(remaining.pop(0)) while len(selected) < final_k and remaining: best_score = -float("inf") best_idx = 0 for i, r in enumerate(remaining): # 相关性分数 relevance = r.score # 与已选择的最大相似度(多样性惩罚) max_sim = max( abs(s.score * r.score) # 近似相似度 for s in selected ) if selected else 0 mmr = lambda_coef * relevance - (1 - lambda_coef) * max_sim if mmr > best_score: best_score = mmr best_idx = i selected.append(remaining.pop(best_idx)) return selected async def find_similar_pairs(self, top_k: int = 20) -> list[dict]: """发现图库中高相似度的图片对(用于去重和聚类)。""" # 获取所有图片 scroll_result = self._client.scroll( collection_name=self.IMAGE_COLLECTION, limit=1000, with_vectors=True, )[0] if len(scroll_result) < 2: return [] vectors = [r.vector for r in scroll_result] ids = [r.id for r in scroll_result] # 批量计算余弦相似度矩阵 import torch mat = torch.tensor(np.array(vectors)) sim_matrix = torch.mm(mat, mat.T) similar_pairs = [] for i in range(len(vectors)): for j in range(i + 1, len(vectors)): sim = float(sim_matrix[i][j]) if sim > 0.95: # 高相似度阈值 similar_pairs.append({ "asset_a": ids[i], "asset_b": ids[j], "similarity": sim, }) similar_pairs.sort(key=lambda x: x["similarity"], reverse=True) return similar_pairs[:top_k] # ── 使用示例 ──────────────────────────────────────────── async def main(): encoder = CLIPEncoder() engine = CrossModalSearchEngine(encoder) await engine.initialize() # 索引图片和文本资产 assets = [ MediaAsset( asset_id="img_001", modality=Modality.IMAGE, content="/path/to/summer_banner_blue.png", tags=["banner", "summer", "blue"], metadata={"season": "summer", "campaign": "七月大促"}, ), MediaAsset( asset_id="txt_001", modality=Modality.TEXT, content="夏日促销活动 banner,蓝色海洋渐变背景,清爽简约风格,产品居中展示", tags=["prompt", "summer"], metadata={"author": "designer_a"}, ), ] await engine.index_assets(assets) # 文本搜图 results = await engine.search( query_text="蓝色渐变科技感 banner", mode=SearchMode.TEXT_TO_IMAGE, filter_tags=["banner"], ) print(f"文本搜图: {len(results)} 条结果") for r in results: print(f" [{r.score:.3f}] {r.asset.content}") # 图搜图 results = await engine.search( query_image_path="/path/to/competitor_banner.png", mode=SearchMode.IMAGE_TO_IMAGE, ) print(f"图搜图: {len(results)} 条结果") # 图搜文本 results = await engine.search( query_image_path="/path/to/reference.png", mode=SearchMode.IMAGE_TO_TEXT, ) print(f"图搜文本: 找到匹配的描述/标签") for r in results: print(f" [{r.score:.3f}] {r.asset.content[:100]}") if __name__ == "__main__": asyncio.run(main())

核心设计决策:

  • 图片和文本分集合存储:虽然它们在同一个向量空间,但分开存有两个好处:一是查询时不需要过滤 modality,提高效率;二是可以做不同的索引参数调优(图片集合可能需要更大的ef_construct)。
  • 单一搜索入口,模式参数化search()函数的mode参数路由到不同检索路径。使用者不需要知道"图片索引叫media_images,文本索引叫media_texts"这些细节。
  • 混合查询的向量融合:文本向量和图片向量直接做加权平均,前提是它们已经 L2 归一化。归一化后才能保证加权融合后不偏离单位球。
  • MMR 去重:检索可能返回多张视觉上几乎一样的图(同一张素材的不同版本、不同尺寸)。MMR 保证结果多样性,让用户看到不同的选择。

四、边界分析与架构权衡

4.1 CLIP 的局限性——它不懂文字

CLIP 擅长的是理解图片的整体视觉风格和语义概念。但它有个显著的弱点:不擅长识别图片中的文字

如果你搜"图片里有'全场5折'这几个字的 banner",CLIP 可能找不到——因为它的训练目标是匹配图文对的整体语义,不是 OCR。对于需要识别图中文字的场景,需要额外的 OCR 处理层(提取图中文字存到元数据,走文本搜索)。

4.2 检索延迟

CLIP 编码一张图片约 10-50ms(取决于硬件),Qdrant 检索 100 万向量约 1-5ms。整体检索延迟 15-55ms,对交互式场景完全够用。

如果图片库达到千万级别,需要用 FAISS 的 IVF + PQ 索引来加速。但绝大多数企业素材库不会超过百万,Qdrant 单机足够。

4.3 可解释性

"相似度 0.87"对于非技术用户完全没意义。你需要解释为什么这两张图相似。

一个实用的方案是用 CLIP 的注意力权重做可视化:高亮查询图片和结果图片中共通的关键区域(色调、构图、主体)。或者用大语言模型(GPT-4 Vision)来生成自然语言解释:"这两张 banner 的相似之处在于都使用了蓝橙对比色调、居中构图、以及圆角卡片式产品展示。"

4.4 中文 CLIP 的选择

标准的 OpenAI CLIP 对中文支持有限。中文场景推荐使用:

  • Chinese-CLIP:OFA-Sys 的中文 CLIP 变体,支持中英文双语。
  • AltCLIP:智源研究院的多语言 CLIP,支持中文和英文。
  • M-CLIP:通过多语言蒸馏增强的 CLIP。

替换只需要改CLIPEncodermodel_name参数,架构其余部分不需要变动。

五、总结

跨模态检索的工程本质是:用 CLIP 把文本和图片映射到同一个向量空间,然后检索就变成了简单的最近邻搜索

三个关键决策:

  1. 分集合存储,但共享编码器:图片和文本用不同的 Qdrant collection 存储,但使用同一个 CLIP 编码器。这样 query 时可以根据模式直接路由到目标集合。
  2. 混合查询是杀手级特性:文本 + 参考图的混合查询,通过加权融合两个向量来实现。α 参数让用户可以在"更像文字描述"和"更像参考图"之间调节。
  3. 不要忽视可解释性:对于创意工作者来说,知道推荐原因比知道分数更重要。预留注意力可视化和 LLM 解释的接口。

这个架构的可扩展性很好。新做视频检索——用 VideoCLIP 替换编码器。新做 3D 模型检索——用 PointCLIP。架构不变,只换编码器。这就是抽象的价值。


下一篇预告:技术博客怎么写才能在保证质量的同时高效交付?聊聊我全职写 10 篇技术文章的流程复盘。