如何用Python+Redis存储张量/数组并实现向量更新与批量相似度检索
解决方案
1. 环境准备
- 安装依赖:
pip install redis numpy - 确保Redis服务为Redis Stack版本(内置Redis Search模块,支持向量检索功能)
2. 初始化Redis连接
import redis import numpy as np # 根据实际Redis配置修改参数 r = redis.Redis(host="localhost", port=6379, db=0, decode_responses=False)
3. 创建向量索引
我们用HASH结构存储每个ID对应的embedding,同时创建FT索引来支持向量检索:
INDEX_NAME = "embedding_index" # 检查索引是否存在,不存在则创建 try: r.ft(INDEX_NAME).info() except redis.ResponseError: # 定义索引结构:指定1024维FLOAT32向量,使用HNSW算法,余弦距离度量 schema = ( redis.search.SchemaField( name="embedding", type=redis.search.VectorField( "embedding", "HNSW", { "TYPE": "FLOAT32", "DIM": 1024, "DISTANCE_METRIC": "COSINE" } ) ), ) # 只索引key前缀为"embedding:"的HASH结构 index_def = redis.search.IndexDefinition(prefix=["embedding:"], index_type=redis.search.IndexType.HASH) r.ft(INDEX_NAME).create_index(schema, definition=index_def)
4. 存储/更新Embedding
将扁平化的1024维向量转为FLOAT32字节数组后存入Redis,同一ID重复调用即可覆盖更新:
def save_or_update_embedding(doc_id, embedding): # 统一转换为FLOAT32类型的numpy数组 if isinstance(embedding, list): embedding = np.array(embedding, dtype=np.float32) elif embedding.dtype != np.float32: embedding = embedding.astype(np.float32) # 存入HASH结构,key格式为embedding:{doc_id} r.hset(f"embedding:{doc_id}", mapping={"embedding": embedding.tobytes()})
调用示例:
# 存储ID为doc_001的1024维向量 save_or_update_embedding("doc_001", np.random.rand(1024)) # 更新同一ID的向量值 save_or_update_embedding("doc_001", np.random.rand(1024))
5. 批量向量检索
对形状为[3,1024]的批量输入向量,逐个执行top-k相似查询:
def batch_search_top_k(input_vectors, top_k=5): results = [] ft_client = r.ft(INDEX_NAME) for vec in input_vectors: # 将输入向量转为FLOAT32字节数组 vec_bytes = vec.astype(np.float32).tobytes() # 构建KNN查询,返回匹配的文档ID并按相似度排序 query = redis.search.Query(f"*=>[KNN {top_k} @embedding $vec]") \ .return_fields("__key") \ .sort_by("__score") \ .paging(0, top_k) # 执行查询,传入向量参数 search_res = ft_client.search(query, query_params={"vec": vec_bytes}) # 提取文档ID(去掉前缀embedding:) top_ids = [doc.id.decode().replace("embedding:", "") for doc in search_res.docs] results.append(top_ids) return results
调用示例:
# 模拟3个1024维输入向量 batch_vecs = np.random.rand(3, 1024) # 获取每个向量的Top-3相似ID top_k_results = batch_search_top_k(batch_vecs, top_k=3) for i, ids in enumerate(top_k_results): print(f"向量{i+1}的Top-3相似ID: {ids}")
关键说明
- 向量必须以FLOAT32字节数组格式存储,这是Redis Search向量字段的要求。
- 索引创建时的
DISTANCE_METRIC可根据需求选择COSINE、L2或IP。 - 查询结果中的
__score为相似度得分(余弦距离场景下,得分越接近1相似度越高)。
内容的提问来源于stack exchange,提问作者ddrong
相关产品推荐
相关产品推荐

