198 lines
6.4 KiB
Python
198 lines
6.4 KiB
Python
"""
|
||
Q&A 自动拆分分块策略
|
||
|
||
使用 LLM 将文档内容自动拆分为问答对。
|
||
每个分段的 content 存储 question,metadata 中存储 answer。
|
||
检索时用 question 做向量匹配,返回 answer 作为上下文。
|
||
"""
|
||
import json
|
||
import logging
|
||
from typing import Dict, Any, List, Optional
|
||
|
||
from .base import BaseChunker, ChunkResult
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# Q&A 拆分的系统提示词
|
||
QA_SYSTEM_PROMPT = """你是一个专业的知识库问答对生成助手。请根据给定的文本内容,生成高质量的问答对(Q&A pairs)。
|
||
|
||
要求:
|
||
1. 问题应该是用户可能会问的自然语言问题
|
||
2. 答案应该准确、完整,直接来源于原文
|
||
3. 每个问答对应该覆盖文本中的一个独立知识点
|
||
4. 问题要具体明确,避免过于宽泛
|
||
5. 答案要简洁但完整,包含必要的上下文
|
||
|
||
请以 JSON 数组格式输出,每个元素包含 question 和 answer 字段:
|
||
```json
|
||
[
|
||
{"question": "问题1", "answer": "答案1"},
|
||
{"question": "问题2", "answer": "答案2"}
|
||
]
|
||
```
|
||
|
||
只输出 JSON 数组,不要输出其他内容。"""
|
||
|
||
|
||
class QAChunker(BaseChunker):
|
||
"""
|
||
Q&A 自动拆分分块器
|
||
|
||
使用 LLM 将文本拆分为问答对。
|
||
需要在初始化时传入 LLM 调用函数。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
chunk_size: int = 500,
|
||
chunk_overlap: int = 50,
|
||
separator: Optional[str] = None,
|
||
llm_caller: Optional[Any] = None,
|
||
):
|
||
super().__init__(chunk_size, chunk_overlap, separator)
|
||
self._llm_caller = llm_caller
|
||
|
||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||
"""
|
||
同步分块(Q&A 模式不支持同步调用,返回空列表)
|
||
请使用 chunk_async 方法。
|
||
"""
|
||
logger.warning("QAChunker.chunk() 不支持同步调用,请使用 chunk_async()")
|
||
return []
|
||
|
||
async def chunk_async(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||
"""
|
||
异步分块:使用 LLM 将文本拆分为 Q&A 对
|
||
|
||
Args:
|
||
text: 原始文本
|
||
metadata: 文档元数据
|
||
|
||
Returns:
|
||
分块结果列表,每个 ChunkResult 的 content 为 question,
|
||
metadata 中包含 answer 和 chunk_mode='qa'
|
||
"""
|
||
if not text or not text.strip():
|
||
return []
|
||
|
||
metadata = metadata or {}
|
||
text = self._clean_text(text)
|
||
|
||
# 如果文本太长,先按段落粗分再逐段生成 Q&A
|
||
max_input_size = self.chunk_size * 8 # LLM 输入上限
|
||
if len(text) > max_input_size:
|
||
segments = self._split_for_qa(text, max_input_size)
|
||
else:
|
||
segments = [text]
|
||
|
||
all_results = []
|
||
position = 0
|
||
|
||
for segment in segments:
|
||
qa_pairs = await self._generate_qa_pairs(segment)
|
||
for qa in qa_pairs:
|
||
question = qa.get('question', '').strip()
|
||
answer = qa.get('answer', '').strip()
|
||
if not question or not answer:
|
||
continue
|
||
|
||
all_results.append(ChunkResult(
|
||
content=question,
|
||
position=position,
|
||
metadata={
|
||
**metadata,
|
||
'answer': answer,
|
||
'chunk_mode': 'qa',
|
||
},
|
||
))
|
||
position += 1
|
||
|
||
logger.info(f"Q&A 拆分完成: {len(all_results)} 个问答对")
|
||
return all_results
|
||
|
||
async def _generate_qa_pairs(self, text: str) -> List[Dict[str, str]]:
|
||
"""调用 LLM 生成 Q&A 对"""
|
||
if not self._llm_caller:
|
||
logger.error("QAChunker: 未配置 LLM 调用函数")
|
||
return []
|
||
|
||
try:
|
||
user_prompt = f"请根据以下文本生成问答对:\n\n{text}"
|
||
|
||
response_text = await self._llm_caller(
|
||
system_prompt=QA_SYSTEM_PROMPT,
|
||
user_prompt=user_prompt,
|
||
)
|
||
|
||
if not response_text:
|
||
return []
|
||
|
||
# 解析 JSON 响应
|
||
return self._parse_qa_response(response_text)
|
||
|
||
except Exception as e:
|
||
logger.error(f"Q&A 生成失败: {e}")
|
||
return []
|
||
|
||
@staticmethod
|
||
def _parse_qa_response(response_text: str) -> List[Dict[str, str]]:
|
||
"""解析 LLM 返回的 Q&A JSON"""
|
||
try:
|
||
# 尝试直接解析
|
||
result = json.loads(response_text)
|
||
if isinstance(result, list):
|
||
return [
|
||
item for item in result
|
||
if isinstance(item, dict) and 'question' in item and 'answer' in item
|
||
]
|
||
except json.JSONDecodeError:
|
||
pass
|
||
|
||
# 尝试提取 JSON 代码块
|
||
import re
|
||
json_match = re.search(r'```(?:json)?\s*\n?(.*?)\n?```', response_text, re.DOTALL)
|
||
if json_match:
|
||
try:
|
||
result = json.loads(json_match.group(1))
|
||
if isinstance(result, list):
|
||
return [
|
||
item for item in result
|
||
if isinstance(item, dict) and 'question' in item and 'answer' in item
|
||
]
|
||
except json.JSONDecodeError:
|
||
pass
|
||
|
||
# 尝试找到 [ ... ] 部分
|
||
bracket_match = re.search(r'\[.*\]', response_text, re.DOTALL)
|
||
if bracket_match:
|
||
try:
|
||
result = json.loads(bracket_match.group(0))
|
||
if isinstance(result, list):
|
||
return [
|
||
item for item in result
|
||
if isinstance(item, dict) and 'question' in item and 'answer' in item
|
||
]
|
||
except json.JSONDecodeError:
|
||
pass
|
||
|
||
logger.warning(f"无法解析 Q&A 响应: {response_text[:200]}")
|
||
return []
|
||
|
||
def _split_for_qa(self, text: str, max_size: int) -> List[str]:
|
||
"""将长文本按段落分割为适合 LLM 处理的片段"""
|
||
paragraphs = text.split('\n\n')
|
||
segments = []
|
||
current = ""
|
||
|
||
for para in paragraphs:
|
||
if current and len(current) + len(para) + 2 > max_size:
|
||
segments.append(current.strip())
|
||
current = para
|
||
else:
|
||
current = current + "\n\n" + para if current else para
|
||
|
||
if current.strip():
|
||
segments.append(current.strip())
|
||
|
||
return segments
|