41 lines
982 B
Python
41 lines
982 B
Python
import faiss
|
|
import numpy as np
|
|
|
|
# 数据归一化函数
|
|
def normalize_vectors(vectors):
|
|
"""对向量进行L2归一化"""
|
|
norms = np.linalg.norm(vectors, axis=1, keepdims=True)
|
|
# 避免除零
|
|
norms = np.where(norms == 0, 1, norms)
|
|
return vectors / norms
|
|
|
|
|
|
data = np.array([[2, 3], [2, 4], [3, 7]], dtype='float32')
|
|
|
|
# 归一化数据
|
|
data_normalized = normalize_vectors(data)
|
|
print("原始数据:")
|
|
print(data)
|
|
print("归一化后数据:")
|
|
print(data_normalized)
|
|
|
|
# 普通索引
|
|
# base_index = faiss.IndexFlatL2(2)
|
|
base_index = faiss.IndexFlatIP(2)
|
|
|
|
# 包一层 IDMap
|
|
index = faiss.IndexIDMap(base_index)
|
|
|
|
# 指定 ID
|
|
ids = np.array([101, 102, 103]) # 自定义 ID
|
|
# index.add_with_ids(data, ids)
|
|
index.add_with_ids(data_normalized, ids)
|
|
|
|
# 查询
|
|
query = np.array([[3, 4.5]], dtype='float32')
|
|
query_normalized = normalize_vectors(query)
|
|
# D, I = index.search(query, k=2)
|
|
D, I = index.search(query_normalized, k=2)
|
|
print(D)
|
|
print(I) # 可能输出 [[101 102]]
|