sales-assistant-py-new/app/chat_qa/classification/intent_test.py

146 lines
4.4 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.

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_categoryvalue 对应 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())