146 lines
4.4 KiB
Python
146 lines
4.4 KiB
Python
import asyncio
|
||
import json
|
||
from pathlib import Path
|
||
from typing import Dict, List
|
||
|
||
from app.repository.milvus.intent_classification_repository import IntentClassificationRepository
|
||
from app.client.embedding_client_manager import embedding_client
|
||
from app.client.milvus_client_manager import milvus_client
|
||
|
||
|
||
QA_JSON_PATH = Path(__file__).parent / "qa.json"
|
||
|
||
|
||
def parse_qa_json(file_path: Path = None) -> Dict[str, List[str]]:
|
||
"""
|
||
解析 qa.json 文件,key 对应 intent_category,value 对应 content
|
||
"""
|
||
if file_path is None:
|
||
file_path = QA_JSON_PATH
|
||
|
||
print(f"读取文件: {file_path}")
|
||
|
||
if not file_path.exists():
|
||
print(f"文件不存在: {file_path}")
|
||
return {}
|
||
|
||
content = file_path.read_text(encoding="utf-8")
|
||
|
||
try:
|
||
data = json.loads(content)
|
||
except json.JSONDecodeError as e:
|
||
print(f"JSON 解析失败: {e}")
|
||
return {}
|
||
|
||
result: Dict[str, List[str]] = {}
|
||
for key, value in data.items():
|
||
if isinstance(value, str):
|
||
result[key] = [value]
|
||
elif isinstance(value, list):
|
||
result[key] = [str(item) for item in value]
|
||
else:
|
||
result[key] = [json.dumps(value, ensure_ascii=False)]
|
||
|
||
print(f"解析成功,共 {len(result)} 个类别")
|
||
return result
|
||
|
||
|
||
async def insert_intent_data(intent_data: Dict[str, List[str]] = None):
|
||
"""
|
||
将意图分类数据入库到Milvus
|
||
"""
|
||
if intent_data is None:
|
||
intent_data = parse_qa_json()
|
||
|
||
if not intent_data:
|
||
print("没有数据可入库")
|
||
return []
|
||
|
||
repo = IntentClassificationRepository()
|
||
milvus_client.init()
|
||
|
||
repo.create_collection()
|
||
|
||
all_data = []
|
||
all_contents = []
|
||
|
||
for category, contents in intent_data.items():
|
||
for content in contents:
|
||
all_data.append({
|
||
"intent_category": category,
|
||
"content": content
|
||
})
|
||
all_contents.append(content)
|
||
|
||
total = len(all_contents)
|
||
print(f"开始向量化 {total} 条内容...")
|
||
|
||
batch_size = 5
|
||
all_vectors = []
|
||
for start in range(0, total, batch_size):
|
||
end = min(start + batch_size, total)
|
||
batch = all_contents[start:end]
|
||
batch_num = start // batch_size + 1
|
||
total_batches = (total + batch_size - 1) // batch_size
|
||
print(f" 向量化批次 {batch_num}/{total_batches}: {len(batch)} 条")
|
||
|
||
vectors = await embedding_client.aembed_documents(batch)
|
||
print(f" 返回 {len(vectors)} 个向量")
|
||
|
||
if len(vectors) != len(batch):
|
||
print(f" 警告: 向量数量不匹配 (期望 {len(batch)}, 实际 {len(vectors)}),尝试逐条向量化")
|
||
for single_text in batch:
|
||
single_vector = await embedding_client.aembed_documents([single_text])
|
||
if single_vector:
|
||
all_vectors.append(single_vector[0])
|
||
print(f" 单条成功: {single_text[:30]}...")
|
||
else:
|
||
print(f" 单条失败,跳过: {single_text[:30]}...")
|
||
else:
|
||
all_vectors.extend(vectors)
|
||
|
||
print(f"向量化完成,共 {len(all_vectors)} 个向量")
|
||
|
||
if len(all_vectors) != len(all_data):
|
||
print(f"警告: 向量数量({len(all_vectors)})与数据数量({len(all_data)})不匹配,仅入库成功的数据")
|
||
all_data = all_data[:len(all_vectors)]
|
||
|
||
for i, item in enumerate(all_data):
|
||
item["content_vector"] = all_vectors[i]
|
||
|
||
try:
|
||
ids = await repo.save_intent_data(all_data)
|
||
print(f"入库成功,插入 {len(ids)} 条记录")
|
||
return ids
|
||
finally:
|
||
await embedding_client.close()
|
||
|
||
|
||
async def query_by_category(category: str):
|
||
"""
|
||
根据意图类别查询
|
||
"""
|
||
repo = IntentClassificationRepository()
|
||
results = await repo.search_by_intent_category(category)
|
||
print(f"\n查询 '{category}' 结果:")
|
||
for i, item in enumerate(results):
|
||
print(f" {i+1}. {item.get('content', '')[:80]}...")
|
||
|
||
|
||
async def main():
|
||
print("=== 读取 qa.json 并入库 ===")
|
||
|
||
intent_data = parse_qa_json()
|
||
|
||
if intent_data:
|
||
print(f"\n解析到的类别 ({len(intent_data)} 个):")
|
||
for category, contents in intent_data.items():
|
||
print(f" {category}: {len(contents)} 条内容")
|
||
for content in contents:
|
||
print(f" -> {content[:80]}...")
|
||
|
||
await insert_intent_data(intent_data)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
asyncio.run(main()) |