Files
agent_test/qdrant_tool.py

200 lines
5.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Qdrant 图片向量数据库工具
功能:上传 vectors.jsonl 中已分析好的向量,用文本查询返回相似度+图片名
依赖pip install qdrant-client openai
"""
import json
import re
import uuid
from qdrant_client import QdrantClient, models
from openai import OpenAI
# ========== 配置区 ==========
QDRANT_HOST = "10.0.0.66"
QDRANT_PORT = 6333
COLLECTION_NAME = "my_collection1"
VECTOR_DIM = 2048 # vectors.jsonl 中向量维度
# 嵌入模型配置(查询时用)
EMBED_API_KEY = "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"
EMBED_BASE_URL = "https://integrate.api.nvidia.com/v1"
EMBED_MODEL = "nvidia/llama-nemotron-embed-vl-1b-v2"
def get_qdrant_client() -> QdrantClient:
return QdrantClient(host=QDRANT_HOST, port=QDRANT_PORT)
def get_text_vector(text: str) -> list[float]:
"""获取文本向量(用于查询)"""
client = OpenAI(api_key=EMBED_API_KEY, base_url=EMBED_BASE_URL)
resp = client.embeddings.create(
input=[text],
model=EMBED_MODEL,
encoding_format="float",
extra_body={"modality": ["text"], "input_type": "query", "truncate": "NONE"}
)
return resp.data[0].embedding
def parse_timestamp(filename: str) -> str:
"""
从文件名解析时间戳
文件名格式: xxx_00m05s_213.jpg → "00:05"
"""
match = re.search(r'(\d{2})m(\d{2})s', filename)
if match:
return f"{match.group(1)}:{match.group(2)}"
return ""
# ========== 创建集合 ==========
def create_collection(
collection_name: str = COLLECTION_NAME,
vector_dim: int = VECTOR_DIM
):
"""创建 Qdrant 集合"""
client = get_qdrant_client()
collections = [c.name for c in client.get_collections().collections]
if collection_name in collections:
print(f"集合 '{collection_name}' 已存在,跳过创建")
return
client.create_collection(
collection_name=collection_name,
vectors_config=models.VectorParams(
size=vector_dim,
distance=models.Distance.COSINE
)
)
print(f"集合 '{collection_name}' 创建成功,维度={vector_dim}")
# ========== 上传向量 ==========
def upload_vectors_jsonl(
jsonl_path: str = "vectors.jsonl",
collection_name: str = COLLECTION_NAME
):
"""
读取 vectors.jsonl 并上传到 Qdrant
jsonl 格式: {"image_path": "/path/to/image.jpg","image_name":"[猫和....01s_042.jpg", "vector": [...]}
"""
client = get_qdrant_client()
points = []
with open(jsonl_path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
item = json.loads(line)
path = item["path"]
vector = item["vector"]
image_name = path.split("/")[-1].split("\\")[-1]
timestamp = parse_timestamp(image_name)
points.append(
models.PointStruct(
id=str(uuid.uuid4()),
vector=vector,
payload={
"image_name": image_name,
"image_path": path,
"timestamp": timestamp,
}
)
)
sp = serch_payload(image_name)
# print(sp)
if sp[0]:
print(f"向量已经上传过,不再处理{path}")
points = []
continue
# 每 500 条批量上传
if len(points) >= 100:
client.upsert(collection_name=collection_name, points=points)
print(f"已上传 {len(points)} 条...")
points = []
if points:
client.upsert(collection_name=collection_name, points=points)
total = client.count(collection_name=collection_name).count
print(f"上传完成,集合 '{collection_name}'{total} 条向量")
# ========== 查询 ==========
def search(
query_text: str,
top_k: int = 5,
collection_name: str = COLLECTION_NAME
):
"""
用文本搜索相似图片
返回: [(相似度, 图片名, 时间戳), ...]
"""
client = get_qdrant_client()
vector = get_text_vector(query_text)
results = client.query_points(
collection_name=collection_name,
query=vector,
limit=top_k
)
output = []
print(f"查询: '{query_text}'\n")
for i, point in enumerate(results.points):
name = point.payload.get("image_name", "")
ts = point.payload.get("timestamp", "")
score = point.score
output.append((score, name, ts))
print(f" [{i + 1}] 相似度={score:.4f} 时间={ts} 文件={name}")
return output
def serch_payload(exact_value):
# exact_value = '[猫和老鼠(五十周年纪念版-国粤英三语)].Tom.And.Jerry.E01.2001.DVDrip.x264.AC3-CMCT_102m13s_210.jpg'
filter_condition = models.Filter(
must=[
models.FieldCondition(
key="image_name",
match=models.MatchValue(value=exact_value)
)
]
)
client = get_qdrant_client()
result = client.scroll(
collection_name=COLLECTION_NAME,
scroll_filter=filter_condition,
limit=10,
with_payload=True
)
return result
# ========== 运行设置(改这里)==========
# 设为 True 的才会执行
DO_CREATE = False # 创建集合
DO_UPLOAD = 1 # 上传 vectors.jsonl
DO_SEARCH = False # 文本搜图片
SEARCH_TEXT = "猫和老鼠" # 搜索内容
SEARCH_TOP_K = 5 # 返回数量
# ======================================
if __name__ == "__main__":
if DO_CREATE:
create_collection()
if DO_UPLOAD:
upload_vectors_jsonl("vectors1.jsonl")
if DO_SEARCH:
search(SEARCH_TEXT, top_k=SEARCH_TOP_K)