77 lines
2.6 KiB
Python
77 lines
2.6 KiB
Python
"""
|
||
读取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\data.txt"
|
||
asyncio.run(save_qa_milvus(file_path)) |