You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.03 21:53:32