在FastAPI中使用pgvector-python高效获取最相似数据的方法
使用pgvector-python批量高效检索余弦相似度最高的向量
核心结论
完全不需要循环执行单条SQL,pgvector-python结合SQLAlchemy可以实现批量查询优化,一次SQL请求就能处理所有新向量的相似度检索。
前置准备:创建向量索引
高效检索的前提是给embedding字段创建余弦相似度索引,推荐使用HNSW索引(适合高维向量,检索性能优于IVFFlat):
CREATE INDEX idx_embedding_hnsw ON embedding USING hnsw (embedding vector_cosine_ops);
如果数据量较小,也可以用IVFFlat索引:
CREATE INDEX idx_embedding_ivfflat ON embedding USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100);
批量检索实现代码
利用PostgreSQL的unnest函数将多个新向量拆分为行,结合SQLAlchemy的distinct_on特性,一次性获取每个新向量的最相似数据:
from sqlalchemy import select, func, unnest from sqlalchemy.orm import Session from pgvector.sqlalchemy import Vector from your_app.models import OcrModel # 替换为你的模型导入路径 def batch_get_top_similar(db: Session, new_vectors: list[list[float]]): # 将新向量转换为pgvector兼容格式,并为每个向量分配唯一标识(用于匹配结果) pg_vectors = [Vector(vec) for vec in new_vectors] query_ids = list(range(len(pg_vectors))) # 使用unnest将向量列表拆分为带ID的临时行 vec_temp_table = func.unnest(pg_vectors, query_ids).alias("temp_vec") # 构造子查询:计算每个新向量与数据库中所有向量的余弦距离,取每个query_id的最小距离(最相似)数据 subquery = ( select( OcrModel.user_id, OcrModel.embedding.cosine_distance(func.cast(vec_temp_table.c.unnest, Vector(1536))).label("distance"), vec_temp_table.c.unnest_2.label("query_id") ) .cross_join(vec_temp_table) .order_by("query_id", "distance") .distinct_on("query_id") # 每个query_id仅保留距离最小的一条记录 ).subquery() # 执行查询并按原向量顺序整理结果 results = db.execute( select(subquery.c.user_id, subquery.c.distance, subquery.c.query_id) .order_by(subquery.c.query_id) ).all() # 余弦相似度 = 1 - 余弦距离,整理成易读格式 return [ { "original_vector_index": res.query_id, "user_id": res.user_id, "cosine_similarity": round(1 - res.distance, 4) } for res in results ]
使用示例
在FastAPI接口中调用该函数:
from fastapi import FastAPI, Depends from sqlalchemy.orm import Session from your_app.database import get_db # 替换为你的数据库会话获取函数 app = FastAPI() @app.post("/batch-similarity") def batch_similarity(data_new: list[list[float]], db: Session = Depends(get_db)): results = batch_get_top_similar(db, data_new) return {"results": results}
关键优势
- 无循环批量处理:一次SQL请求完成所有向量的检索,避免多次数据库连接开销
- 利用数据库优化:依赖PostgreSQL的
unnest和distinct_on特性,结合pgvector的原生向量运算,性能远高于循环单查 - 可扩展性强:如果需要获取每个向量的Top-N相似数据,只需将
distinct_on替换为窗口函数:from sqlalchemy import over, row_number subquery = ( select( OcrModel.user_id, OcrModel.embedding.cosine_distance(func.cast(vec_temp_table.c.unnest, Vector(1536))).label("distance"), vec_temp_table.c.unnest_2.label("query_id"), row_number().over(partition_by="query_id", order_by="distance").label("rank") ) .cross_join(vec_temp_table) ).subquery() results = db.execute( select(subquery.c.user_id, subquery.c.distance, subquery.c.query_id) .where(subquery.c.rank <= 3) # 取Top3 .order_by(subquery.c.query_id, subquery.c.rank) ).all()
内容的提问来源于stack exchange,提问作者CutePoison
相关产品推荐
相关产品推荐

