Files
agent_test/cha_xianglaing.py

239 lines
7.7 KiB
Python
Raw 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.
import os
import sys
import json
import base64
import time
from pathlib import Path
from multiprocessing import Pool
# ========== 配置区(直接改这里)==========
# 图片目录
IMAGE_DIR = "./shipin/"
# API Keys填多个就多进程并行
API_KEYS = [
"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",
"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",
"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"
]
collection_name = "my_collection1"
# API 地址
BASE_URL = "https://integrate.api.nvidia.com/v1"
# 模型名称
MODEL = "nvidia/llama-nemotron-embed-vl-1b-v2"
# 输出文件
OUTPUT = "vectors1.jsonl"
# 同一 key 请求间隔(秒),防限流
DELAY = 0.1
# ========== 配置结束 ==========
IMAGE_EXTENSIONS = {'.jpg', '.jpeg', '.png', '.bmp', '.webp'}
# ========== 独立函数(供 weixin_bot 等外部调用)==========
def get_client(api_key=None, base_url=None):
"""获取 OpenAI 客户端(单 key 模式,兼容旧接口)"""
from openai import OpenAI
return OpenAI(
api_key=api_key or API_KEYS[0],
base_url=base_url or BASE_URL
)
def encode_image_base64(image_path: str) -> str:
"""将图片编码为 base64"""
with open(image_path, "rb") as f:
return base64.b64encode(f.read()).decode("utf-8")
def get_image_vector(image_path: str, client=None) -> list[float]:
if client is None:
client = get_client()
b64, mime = encode_image_base64(image_path) # 确保你的函数能同时返回这些值
url = f"data:{mime};base64,{b64}"
input_payload = [
{
"content": [
{"type": "text", "text": "What is in this image?"},
{"type": "image_url", "image_url": {"url": url}}
]
}
]
resp = client.embeddings.create(
input=input_payload, # 使用新格式
model=MODEL,
encoding_format="float",
extra_body={"input_type": "passage","truncate": "NONE", "modality": ["image"]}
)
return resp.data[0].embedding
def get_text_vector(text: str, client=None) -> list[float]:
"""获取文本的向量(外部可直接调用)"""
if client is None:
client = get_client()
resp = client.embeddings.create(
input=[text],
model=MODEL,
encoding_format="float",
extra_body={"modality": ["text"], "input_type": "query", "truncate": "NONE"}
)
return resp.data[0].embedding
def search1(query_text: str, top_k=5, qdrant_host="10.0.0.66", qdrant_port=6333):
"""在 Qdrant 中搜索(外部可直接调用)"""
from qdrant_client import QdrantClient
vector = get_text_vector(query_text)
qdrant = QdrantClient(host=qdrant_host, port=qdrant_port)
results = qdrant.query_points(
collection_name=collection_name,
query=vector,
limit=top_k
)
return results.points
# ========== 批量向量化(多 key 多进程)==========
def scan_images(root_path: str) -> list[str]:
"""扫描目录下所有图片文件"""
images = []
root = Path(root_path)
if root.is_file() and root.suffix.lower() in IMAGE_EXTENSIONS:
return [str(root.resolve())]
for p in sorted(root.rglob("*")):
if p.is_file() and p.suffix.lower() in IMAGE_EXTENSIONS:
images.append(str(p.resolve()))
return images
def vectorize(image_path: str, api_key: str, base_url: str, model: str) -> dict:
from openai import OpenAI
import mimetypes
import base64
try:
client = OpenAI(api_key=api_key, base_url=base_url)
# 获取 base64 和 MIME 类型
with open(image_path, "rb") as f:
b64 = base64.b64encode(f.read()).decode("utf-8")
mime, _ = mimetypes.guess_type(image_path)
if not mime:
mime = "image/jpeg"
uri = f"data:{mime};base64,{b64}"
resp = client.embeddings.create(
input=[uri],
model=model,
encoding_format="float",
# !!! 关键修改:把 "query" 改为 "passage" !!!
extra_body={"input_type": "passage", "truncate": "NONE", "modality": ["image"]}
)
return {"path": image_path, "vector": resp.data[0].embedding, "status": "ok"}
except Exception as e:
return {"path": image_path, "vector": None, "status": "error", "error": str(e)}
def _worker(args) -> dict:
"""多进程 worker返回结果同时携带进程信息"""
image_path, api_key, base_url, model, delay = args
if delay > 0:
time.sleep(delay)
pid = os.getpid()
# 也可以使用 current_process().name
result = vectorize(image_path, api_key, base_url, model)
result["pid"] = pid
result["api_key_prefix"] = api_key[:10] + "..." # 只显示前几位
return result
def load_existing(output_path: str) -> dict:
"""加载已有的向量 JSONL每行一个 JSON 对象)"""
result = {}
if os.path.exists(output_path):
with open(output_path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
item = json.loads(line)
result[item["path"]] = item
return result
def append_vectors(new_items: list[dict], output_path: str):
"""追加写入新向量到 JSONL 文件(每行一个 JSON 对象)"""
with open(output_path, "a", encoding="utf-8") as f:
for item in new_items:
f.write(json.dumps(item, ensure_ascii=False) + "\n")
def main():
images = scan_images(IMAGE_DIR)
if not images:
print(f"未在 {IMAGE_DIR} 中找到图片文件")
sys.exit(1)
existing = load_existing(OUTPUT)
todo = [img for img in images if img not in existing]
skipped = len(images) - len(todo)
if not todo:
print(f"所有 {len(images)} 张图片已存在于 {OUTPUT},无需处理")
return
jobs = len(API_KEYS)
print(f"找到 {len(images)} 张图片,待处理 {len(todo)},跳过 {skipped}")
print(f"使用 {jobs} 个 API Key并行处理")
# 构建任务:轮询分配 key同一 key 的请求错开 delay
tasks = []
for i, img in enumerate(todo):
key = API_KEYS[i % len(API_KEYS)]
delay = DELAY
# print(delay)
tasks.append((img, key, BASE_URL, MODEL, delay))
done = 0
errors = 0
total = len(existing)
with Pool(jobs) as pool:
for result in pool.imap_unordered(_worker, tasks):
if result["status"] == "ok":
done += 1
append_vectors([{"path": result["path"], "vector": result["vector"]}], OUTPUT)
total += 1
# 打印详细信息进程ID + 文件路径
print(f"[PID {result['pid']}] 已处理: {os.path.basename(result['path'])}")
if done % 50 == 0:
print(f"[进度] 已完成 {done}/{len(todo)},错误 {errors},总计 {total}")
else:
errors += 1
print(f"[错误][PID {result.get('pid', '?')}] {os.path.basename(result['path'])}: {result['error']}")
print(f"\n处理完毕: 完成 {done}, 跳过 {skipped}, 错误 {errors}")
print(f"向量已追加到 {OUTPUT},共 {total} 条记录")
if __name__ == "__main__":
main()
# from qdrant_client import QdrantClient, models
# client = QdrantClient(host="10.0.0.66", port=6333)
# collection_name = "my_collection1" #集合名称
# client.create_collection(
# collection_name=collection_name,
# vectors_config=models.VectorParams(
# size=2048, # 向量维度128维
# distance=models.Distance.COSINE # 距离计算方法,这里用余弦相似度
# )
# )
# client.delete_collection(collection_name="my_collection")