sales-assistant-py-new/app/chat_qa/deduplication_qa/message_qa.py

77 lines
2.6 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.

"""
读取d_qa生成的qa数据入库到向量库qa表中
"""
import ast
import asyncio
import json
from pathlib import Path
from typing import List
from app.client.embedding_client_manager import embedding_client
from app.client.milvus_client_manager import milvus_client
from app.repository.milvus.message_qa_repository import QARepository
# 字符串转 字典
def safe_parse_json(raw_data):
# 本身就是字典,直接返回
if isinstance(raw_data, dict):
return raw_data
if not isinstance(raw_data, str):
return None
s = raw_data.strip()
if not s:
return None
try:
return ast.literal_eval(s)
except (ValueError, SyntaxError):
try:
return json.loads(s)
except json.JSONDecodeError as e:
print(f"警告: JSON解析失败: {e},跳过")
return None
async def read_file(file_path:str)->List[dict[str, str]]:
qa_file = Path(file_path)
qa_list = []
with open(qa_file, 'r', encoding='utf-8') as f:
for line_num, line in enumerate(f, 1):
line = line.strip()
if not line:
continue
qa_list.append(safe_parse_json( line))
return qa_list
async def save_qa_milvus(file_path:str):
try :
embedding_client.init()
milvus_client.init()
qa_milvus_repository = QARepository()
qa_milvus_repository.create_collection()
final_qas: List[dict[str, str]] = await read_file(file_path)
question_list: List[str] = [qa["question"] for qa in final_qas]
embedding_batch_size = 10
embedding_ed_all: list[list[float]] = []
for i in range(0, len(question_list), embedding_batch_size):
embedding_part_texts = question_list[i:i + embedding_batch_size]
embedding_eds = await embedding_client.aembed_documents(embedding_part_texts)
embedding_ed_all.extend(embedding_eds)
# 构建对象列表
qa_list = []
for i, qa in enumerate(final_qas):
qa_list.append({
"question": qa["question"],
"answer": qa.get("answer", ""),
"question_dense_vector": embedding_ed_all[i] if i < len(embedding_ed_all) else [],
"parent_id": qa.get("parent_id", 0)
})
await qa_milvus_repository.save_qa_to_milvus(qa_list)
except Exception as e:
print(e)
finally:
await embedding_client.close()
milvus_client.close()
if __name__ == '__main__':
file_path=r"E:\qa_output\test.txt"
asyncio.run(save_qa_milvus(file_path))