Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,45 @@
|
||||
"""
|
||||
文档分块策略
|
||||
"""
|
||||
from .base import BaseChunker, ChunkResult
|
||||
from .recursive import RecursiveChunker
|
||||
from .markdown import MarkdownChunker
|
||||
from .fixed import FixedChunker
|
||||
from .qa_chunker import QAChunker
|
||||
from .sentence import SentenceChunker
|
||||
from .auto import AutoChunker
|
||||
|
||||
__all__ = [
|
||||
'BaseChunker',
|
||||
'ChunkResult',
|
||||
'RecursiveChunker',
|
||||
'MarkdownChunker',
|
||||
'FixedChunker',
|
||||
'QAChunker',
|
||||
'SentenceChunker',
|
||||
'AutoChunker',
|
||||
'get_chunker',
|
||||
]
|
||||
|
||||
|
||||
def get_chunker(strategy: str, chunk_size: int = 500, chunk_overlap: int = 50, separator: str = None, **kwargs) -> BaseChunker:
|
||||
"""
|
||||
根据策略名称获取分块器实例
|
||||
|
||||
Args:
|
||||
strategy: 分块策略名称(recursive/markdown/fixed/qa/sentence/auto)
|
||||
chunk_size: 分块大小
|
||||
chunk_overlap: 分块重叠
|
||||
separator: 自定义分隔符
|
||||
**kwargs: 额外参数(如 QAChunker 的 llm_caller)
|
||||
"""
|
||||
chunkers = {
|
||||
'recursive': RecursiveChunker,
|
||||
'markdown': MarkdownChunker,
|
||||
'fixed': FixedChunker,
|
||||
'qa': QAChunker,
|
||||
'sentence': SentenceChunker,
|
||||
'auto': AutoChunker,
|
||||
}
|
||||
chunker_cls = chunkers.get(strategy, RecursiveChunker)
|
||||
return chunker_cls(chunk_size=chunk_size, chunk_overlap=chunk_overlap, separator=separator, **kwargs)
|
||||
@@ -0,0 +1,90 @@
|
||||
"""
|
||||
自动分块策略
|
||||
|
||||
根据文件类型自动选择最佳分块器。
|
||||
参考 Dify 的 auto 分块模式。
|
||||
"""
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 文件类型 → 推荐分块策略
|
||||
FILE_TYPE_STRATEGY_MAP = {
|
||||
# Markdown 文件使用 Markdown 分块器
|
||||
'md': 'markdown',
|
||||
'markdown': 'markdown',
|
||||
# 代码文件使用按句子分块(按行/语句边界)
|
||||
'py': 'sentence',
|
||||
'js': 'sentence',
|
||||
'ts': 'sentence',
|
||||
'java': 'sentence',
|
||||
'go': 'sentence',
|
||||
'rs': 'sentence',
|
||||
'c': 'sentence',
|
||||
'cpp': 'sentence',
|
||||
'h': 'sentence',
|
||||
# 纯文本使用按句子分块
|
||||
'txt': 'sentence',
|
||||
# CSV/Excel 使用固定大小(表格数据按行分割更合理)
|
||||
'csv': 'fixed',
|
||||
'xlsx': 'fixed',
|
||||
'xls': 'fixed',
|
||||
# HTML 使用 Markdown 分块器(HTML 结构类似)
|
||||
'html': 'markdown',
|
||||
'htm': 'markdown',
|
||||
# 其他文档类型使用递归分块
|
||||
'pdf': 'recursive',
|
||||
'docx': 'recursive',
|
||||
'doc': 'recursive',
|
||||
'pptx': 'recursive',
|
||||
'ppt': 'recursive',
|
||||
}
|
||||
|
||||
|
||||
class AutoChunker(BaseChunker):
|
||||
"""
|
||||
自动分块器
|
||||
|
||||
根据文档的文件类型自动选择最佳分块策略。
|
||||
metadata 中需要包含 'file_type' 字段。
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""自动选择分块策略并执行"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
metadata = metadata or {}
|
||||
file_type = metadata.get('file_type', '').lower().lstrip('.')
|
||||
|
||||
# 根据文件类型选择策略
|
||||
strategy = FILE_TYPE_STRATEGY_MAP.get(file_type, 'recursive')
|
||||
|
||||
logger.info(f"AutoChunker: file_type={file_type} → strategy={strategy}")
|
||||
|
||||
# 动态创建对应的分块器
|
||||
chunker = self._get_chunker(strategy)
|
||||
return chunker.chunk(text, metadata)
|
||||
|
||||
def _get_chunker(self, strategy: str) -> BaseChunker:
|
||||
"""获取对应策略的分块器实例"""
|
||||
from .recursive import RecursiveChunker
|
||||
from .markdown import MarkdownChunker
|
||||
from .fixed import FixedChunker
|
||||
from .sentence import SentenceChunker
|
||||
|
||||
chunkers = {
|
||||
'recursive': RecursiveChunker,
|
||||
'markdown': MarkdownChunker,
|
||||
'fixed': FixedChunker,
|
||||
'sentence': SentenceChunker,
|
||||
}
|
||||
cls = chunkers.get(strategy, RecursiveChunker)
|
||||
return cls(
|
||||
chunk_size=self.chunk_size,
|
||||
chunk_overlap=self.chunk_overlap,
|
||||
separator=self.separator,
|
||||
)
|
||||
@@ -0,0 +1,92 @@
|
||||
"""
|
||||
分块策略基类
|
||||
"""
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChunkResult:
|
||||
"""分块结果"""
|
||||
content: str
|
||||
position: int = 0
|
||||
char_count: int = 0
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
if not self.char_count:
|
||||
self.char_count = len(self.content)
|
||||
|
||||
|
||||
class BaseChunker(ABC):
|
||||
"""
|
||||
分块策略基类
|
||||
|
||||
所有分块策略必须继承此类并实现 chunk 方法
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chunk_size: int = 500,
|
||||
chunk_overlap: int = 50,
|
||||
separator: Optional[str] = None,
|
||||
):
|
||||
self.chunk_size = chunk_size
|
||||
self.chunk_overlap = chunk_overlap
|
||||
self.separator = separator
|
||||
|
||||
@abstractmethod
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""
|
||||
将文本分块
|
||||
|
||||
Args:
|
||||
text: 原始文本
|
||||
metadata: 文档元数据
|
||||
|
||||
Returns:
|
||||
分块结果列表
|
||||
"""
|
||||
pass
|
||||
|
||||
def _clean_text(self, text: str) -> str:
|
||||
"""清理文本:去除多余空白"""
|
||||
import re
|
||||
# 合并连续空行为单个空行
|
||||
text = re.sub(r'\n{3,}', '\n\n', text)
|
||||
# 去除行尾空白
|
||||
text = '\n'.join(line.rstrip() for line in text.split('\n'))
|
||||
return text.strip()
|
||||
|
||||
def _merge_small_chunks(self, chunks: List[str], min_size: int = 50) -> List[str]:
|
||||
"""合并过小的分块"""
|
||||
if not chunks:
|
||||
return []
|
||||
|
||||
merged = []
|
||||
buffer = ""
|
||||
|
||||
for chunk in chunks:
|
||||
if not chunk.strip():
|
||||
continue
|
||||
if buffer and len(buffer) + len(chunk) <= self.chunk_size:
|
||||
buffer = buffer + "\n" + chunk
|
||||
elif buffer and len(buffer) < min_size:
|
||||
buffer = buffer + "\n" + chunk
|
||||
else:
|
||||
if buffer:
|
||||
merged.append(buffer)
|
||||
buffer = chunk
|
||||
|
||||
if buffer:
|
||||
# 最后一个 buffer 如果太小,合并到前一个
|
||||
if merged and len(buffer) < min_size:
|
||||
merged[-1] = merged[-1] + "\n" + buffer
|
||||
else:
|
||||
merged.append(buffer)
|
||||
|
||||
return merged
|
||||
@@ -0,0 +1,57 @@
|
||||
"""
|
||||
固定大小分块策略
|
||||
|
||||
按固定字符数分割文本,最简单的分块方式
|
||||
"""
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FixedChunker(BaseChunker):
|
||||
"""
|
||||
固定大小分块器
|
||||
|
||||
按固定字符数分割文本,相邻分块之间有 overlap 重叠
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""固定大小分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
if len(text) <= self.chunk_size:
|
||||
return [ChunkResult(
|
||||
content=text,
|
||||
position=0,
|
||||
metadata={**metadata},
|
||||
)]
|
||||
|
||||
chunks = []
|
||||
start = 0
|
||||
position = 0
|
||||
step = self.chunk_size - self.chunk_overlap
|
||||
|
||||
while start < len(text):
|
||||
end = min(start + self.chunk_size, len(text))
|
||||
chunk_text = text[start:end].strip()
|
||||
|
||||
if chunk_text:
|
||||
chunks.append(ChunkResult(
|
||||
content=chunk_text,
|
||||
position=position,
|
||||
metadata={**metadata},
|
||||
))
|
||||
position += 1
|
||||
|
||||
start += step
|
||||
if step <= 0:
|
||||
break
|
||||
|
||||
return chunks
|
||||
@@ -0,0 +1,172 @@
|
||||
"""
|
||||
Markdown 结构化分块策略
|
||||
|
||||
按 Markdown 标题层级分割文档,保留文档结构信息
|
||||
"""
|
||||
import re
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MarkdownChunker(BaseChunker):
|
||||
"""
|
||||
Markdown 分块器
|
||||
|
||||
按标题层级分割 Markdown 文档,每个标题下的内容作为一个分块
|
||||
如果单个标题下的内容超过 chunk_size,则使用递归分割
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""按 Markdown 标题分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
# 按标题分割
|
||||
sections = self._split_by_headers(text)
|
||||
|
||||
# 处理每个 section
|
||||
raw_chunks = []
|
||||
for section in sections:
|
||||
header = section.get('header', '')
|
||||
content = section.get('content', '')
|
||||
level = section.get('level', 0)
|
||||
|
||||
if not content.strip():
|
||||
continue
|
||||
|
||||
# 组合标题和内容
|
||||
full_text = f"{header}\n{content}" if header else content
|
||||
|
||||
if len(full_text) <= self.chunk_size:
|
||||
raw_chunks.append({
|
||||
'content': full_text.strip(),
|
||||
'metadata': {
|
||||
**metadata,
|
||||
'header': header,
|
||||
'header_level': level,
|
||||
}
|
||||
})
|
||||
else:
|
||||
# 内容超长,递归分割
|
||||
sub_chunks = self._split_long_section(content, header)
|
||||
for i, sub in enumerate(sub_chunks):
|
||||
raw_chunks.append({
|
||||
'content': sub.strip(),
|
||||
'metadata': {
|
||||
**metadata,
|
||||
'header': header,
|
||||
'header_level': level,
|
||||
'sub_chunk': i,
|
||||
}
|
||||
})
|
||||
|
||||
# 合并过小的分块
|
||||
merged = self._merge_small_section_chunks(raw_chunks)
|
||||
|
||||
# 构建结果
|
||||
results = []
|
||||
for i, item in enumerate(merged):
|
||||
if item['content'].strip():
|
||||
results.append(ChunkResult(
|
||||
content=item['content'],
|
||||
position=i,
|
||||
metadata=item.get('metadata', {}),
|
||||
))
|
||||
|
||||
return results
|
||||
|
||||
def _split_by_headers(self, text: str) -> List[Dict[str, Any]]:
|
||||
"""按 Markdown 标题分割"""
|
||||
# 匹配 Markdown 标题: # Title, ## Title, ### Title 等
|
||||
header_pattern = re.compile(r'^(#{1,6})\s+(.+)$', re.MULTILINE)
|
||||
|
||||
sections = []
|
||||
last_end = 0
|
||||
last_header = ''
|
||||
last_level = 0
|
||||
|
||||
for match in header_pattern.finditer(text):
|
||||
# 保存前一个 section 的内容
|
||||
if last_end > 0 or match.start() > 0:
|
||||
content = text[last_end:match.start()]
|
||||
if content.strip() or last_header:
|
||||
sections.append({
|
||||
'header': last_header,
|
||||
'content': content.strip(),
|
||||
'level': last_level,
|
||||
})
|
||||
|
||||
last_header = match.group(0)
|
||||
last_level = len(match.group(1))
|
||||
last_end = match.end()
|
||||
|
||||
# 最后一个 section
|
||||
remaining = text[last_end:]
|
||||
if remaining.strip() or last_header:
|
||||
sections.append({
|
||||
'header': last_header,
|
||||
'content': remaining.strip(),
|
||||
'level': last_level,
|
||||
})
|
||||
|
||||
# 如果没有找到任何标题,整个文本作为一个 section
|
||||
if not sections:
|
||||
sections.append({
|
||||
'header': '',
|
||||
'content': text.strip(),
|
||||
'level': 0,
|
||||
})
|
||||
|
||||
return sections
|
||||
|
||||
def _split_long_section(self, content: str, header: str = '') -> List[str]:
|
||||
"""分割超长的 section 内容"""
|
||||
from .recursive import RecursiveChunker
|
||||
|
||||
chunker = RecursiveChunker(
|
||||
chunk_size=self.chunk_size,
|
||||
chunk_overlap=self.chunk_overlap,
|
||||
)
|
||||
results = chunker.chunk(content)
|
||||
|
||||
chunks = []
|
||||
for i, result in enumerate(results):
|
||||
# 第一个分块带上标题
|
||||
if i == 0 and header:
|
||||
chunks.append(f"{header}\n{result.content}")
|
||||
else:
|
||||
chunks.append(result.content)
|
||||
|
||||
return chunks if chunks else [content]
|
||||
|
||||
def _merge_small_section_chunks(self, chunks: List[Dict], min_size: int = 80) -> List[Dict]:
|
||||
"""合并过小的 section 分块"""
|
||||
if not chunks:
|
||||
return []
|
||||
|
||||
merged = []
|
||||
buffer = None
|
||||
|
||||
for chunk in chunks:
|
||||
if buffer is None:
|
||||
buffer = chunk
|
||||
elif len(buffer['content']) < min_size and len(buffer['content']) + len(chunk['content']) <= self.chunk_size:
|
||||
buffer['content'] = buffer['content'] + "\n\n" + chunk['content']
|
||||
else:
|
||||
merged.append(buffer)
|
||||
buffer = chunk
|
||||
|
||||
if buffer:
|
||||
if merged and len(buffer['content']) < min_size:
|
||||
merged[-1]['content'] = merged[-1]['content'] + "\n\n" + buffer['content']
|
||||
else:
|
||||
merged.append(buffer)
|
||||
|
||||
return merged
|
||||
@@ -0,0 +1,197 @@
|
||||
"""
|
||||
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
|
||||
@@ -0,0 +1,160 @@
|
||||
"""
|
||||
递归字符分块策略
|
||||
|
||||
最常用的分块策略,按照分隔符层级递归分割文本
|
||||
优先按段落 → 句子 → 字符的顺序分割
|
||||
"""
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 默认分隔符层级(从大到小)
|
||||
DEFAULT_SEPARATORS = [
|
||||
"\n\n", # 段落
|
||||
"\n", # 换行
|
||||
"。", # 中文句号
|
||||
"!", # 中文感叹号
|
||||
"?", # 中文问号
|
||||
";", # 中文分号
|
||||
". ", # 英文句号
|
||||
"! ", # 英文感叹号
|
||||
"? ", # 英文问号
|
||||
"; ", # 英文分号
|
||||
",", # 中文逗号
|
||||
", ", # 英文逗号
|
||||
" ", # 空格
|
||||
"", # 逐字符
|
||||
]
|
||||
|
||||
|
||||
class RecursiveChunker(BaseChunker):
|
||||
"""
|
||||
递归字符分块器
|
||||
|
||||
按分隔符层级递归分割文本,确保每个分块不超过 chunk_size,
|
||||
相邻分块之间有 chunk_overlap 的重叠
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chunk_size: int = 500,
|
||||
chunk_overlap: int = 50,
|
||||
separator: Optional[str] = None,
|
||||
separators: Optional[List[str]] = None,
|
||||
):
|
||||
super().__init__(chunk_size, chunk_overlap, separator)
|
||||
if separator:
|
||||
self.separators = [separator] + DEFAULT_SEPARATORS
|
||||
elif separators:
|
||||
self.separators = separators
|
||||
else:
|
||||
self.separators = DEFAULT_SEPARATORS
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""递归分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
# 递归分割
|
||||
raw_chunks = self._recursive_split(text, self.separators)
|
||||
|
||||
# 合并过小的分块
|
||||
raw_chunks = self._merge_small_chunks(raw_chunks)
|
||||
|
||||
# 添加重叠
|
||||
chunks_with_overlap = self._add_overlap(raw_chunks)
|
||||
|
||||
# 构建结果
|
||||
results = []
|
||||
for i, content in enumerate(chunks_with_overlap):
|
||||
if content.strip():
|
||||
results.append(ChunkResult(
|
||||
content=content.strip(),
|
||||
position=i,
|
||||
metadata={**metadata},
|
||||
))
|
||||
|
||||
return results
|
||||
|
||||
def _recursive_split(self, text: str, separators: List[str]) -> List[str]:
|
||||
"""递归分割文本"""
|
||||
if len(text) <= self.chunk_size:
|
||||
return [text] if text.strip() else []
|
||||
|
||||
# 找到合适的分隔符
|
||||
separator = ""
|
||||
for sep in separators:
|
||||
if sep == "":
|
||||
separator = sep
|
||||
break
|
||||
if sep in text:
|
||||
separator = sep
|
||||
break
|
||||
|
||||
# 按分隔符分割
|
||||
if separator:
|
||||
splits = text.split(separator)
|
||||
else:
|
||||
# 逐字符分割
|
||||
splits = list(text)
|
||||
|
||||
# 合并分割结果,确保不超过 chunk_size
|
||||
chunks = []
|
||||
current = ""
|
||||
|
||||
for split in splits:
|
||||
piece = split if not separator else split
|
||||
test_piece = current + separator + piece if current else piece
|
||||
|
||||
if len(test_piece) <= self.chunk_size:
|
||||
current = test_piece
|
||||
else:
|
||||
if current:
|
||||
chunks.append(current)
|
||||
# 如果单个片段超过 chunk_size,递归处理
|
||||
if len(piece) > self.chunk_size:
|
||||
remaining_separators = separators[separators.index(separator) + 1:] if separator in separators else separators[1:]
|
||||
if remaining_separators:
|
||||
sub_chunks = self._recursive_split(piece, remaining_separators)
|
||||
chunks.extend(sub_chunks)
|
||||
current = ""
|
||||
else:
|
||||
# 没有更小的分隔符了,强制截断
|
||||
for j in range(0, len(piece), self.chunk_size):
|
||||
chunks.append(piece[j:j + self.chunk_size])
|
||||
current = ""
|
||||
else:
|
||||
current = piece
|
||||
|
||||
if current:
|
||||
chunks.append(current)
|
||||
|
||||
return chunks
|
||||
|
||||
def _add_overlap(self, chunks: List[str]) -> List[str]:
|
||||
"""为相邻分块添加重叠"""
|
||||
if self.chunk_overlap <= 0 or len(chunks) <= 1:
|
||||
return chunks
|
||||
|
||||
result = []
|
||||
for i, chunk in enumerate(chunks):
|
||||
if i == 0:
|
||||
result.append(chunk)
|
||||
else:
|
||||
# 从前一个分块的末尾取 overlap 字符作为前缀
|
||||
prev = chunks[i - 1]
|
||||
overlap_text = prev[-self.chunk_overlap:] if len(prev) > self.chunk_overlap else prev
|
||||
# 确保合并后不超过 chunk_size 太多
|
||||
combined = overlap_text + "\n" + chunk
|
||||
if len(combined) <= self.chunk_size * 1.2:
|
||||
result.append(combined)
|
||||
else:
|
||||
result.append(chunk)
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,105 @@
|
||||
"""
|
||||
按句子分块策略
|
||||
|
||||
按句号/问号/感叹号等句子边界分割文本,
|
||||
然后将小句子合并到不超过 chunk_size 的分块中。
|
||||
参考 Dify 的 sentence 分块模式。
|
||||
"""
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 句子分隔符正则(中英文句号/问号/感叹号/分号)
|
||||
SENTENCE_PATTERN = re.compile(
|
||||
r'(?<=[。!?;.!?;])\s*'
|
||||
)
|
||||
|
||||
|
||||
class SentenceChunker(BaseChunker):
|
||||
"""
|
||||
按句子分块器
|
||||
|
||||
先按句子边界分割文本,再将相邻句子合并为不超过 chunk_size 的分块。
|
||||
保证每个分块都是完整句子的组合,不会在句子中间截断。
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""按句子分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
# 按句子边界分割
|
||||
sentences = SENTENCE_PATTERN.split(text)
|
||||
sentences = [s.strip() for s in sentences if s.strip()]
|
||||
|
||||
if not sentences:
|
||||
return [ChunkResult(content=text, position=0, metadata={**metadata})]
|
||||
|
||||
# 合并句子为分块(不超过 chunk_size)
|
||||
chunks = []
|
||||
current = ""
|
||||
position = 0
|
||||
|
||||
for sentence in sentences:
|
||||
# 如果单个句子就超过 chunk_size,强制作为独立分块
|
||||
if len(sentence) > self.chunk_size:
|
||||
if current:
|
||||
chunks.append(current)
|
||||
current = ""
|
||||
chunks.append(sentence)
|
||||
continue
|
||||
|
||||
test = current + sentence if not current else current + " " + sentence
|
||||
if len(test) <= self.chunk_size:
|
||||
current = test
|
||||
else:
|
||||
if current:
|
||||
chunks.append(current)
|
||||
current = sentence
|
||||
|
||||
if current:
|
||||
chunks.append(current)
|
||||
|
||||
# 合并过小的分块
|
||||
chunks = self._merge_small_chunks(chunks)
|
||||
|
||||
# 添加重叠
|
||||
if self.chunk_overlap > 0 and len(chunks) > 1:
|
||||
chunks = self._add_sentence_overlap(chunks)
|
||||
|
||||
# 构建结果
|
||||
results = []
|
||||
for i, content in enumerate(chunks):
|
||||
if content.strip():
|
||||
results.append(ChunkResult(
|
||||
content=content.strip(),
|
||||
position=i,
|
||||
metadata={**metadata},
|
||||
))
|
||||
|
||||
return results
|
||||
|
||||
def _add_sentence_overlap(self, chunks: List[str]) -> List[str]:
|
||||
"""为相邻分块添加句子级重叠"""
|
||||
result = [chunks[0]]
|
||||
for i in range(1, len(chunks)):
|
||||
prev = chunks[i - 1]
|
||||
# 从前一个分块取最后一个句子作为重叠
|
||||
prev_sentences = SENTENCE_PATTERN.split(prev)
|
||||
prev_sentences = [s.strip() for s in prev_sentences if s.strip()]
|
||||
if prev_sentences:
|
||||
overlap = prev_sentences[-1]
|
||||
if len(overlap) <= self.chunk_overlap:
|
||||
combined = overlap + " " + chunks[i]
|
||||
if len(combined) <= self.chunk_size * 1.2:
|
||||
result.append(combined)
|
||||
continue
|
||||
result.append(chunks[i])
|
||||
return result
|
||||
Reference in New Issue
Block a user