Build lightweight AI agent admin
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user