text_splitter.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347
  1. """文档三级分块服务(接收已解析文本,不直接读取文件)"""
  2. import re
  3. import unicodedata
  4. from typing import Dict, List
  5. from backend.indexing.semantic_chunker import SemanticChunker
  6. # 编译非打印 C0/C1 控制字符的正则(保留常规排版字:\t, \n, \r)
  7. _CONTROL_CHAR_RE = re.compile(r"[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]")
  8. # 编译零宽字符和不可见格式化控制字符(零宽空白、BOM 标记、左右强排标志等)
  9. _INVISIBLE_CHAR_RE = re.compile(r"[\u200b-\u200d\ufeff\u200f\u202a-\u202e]")
  10. def sanitize_text(text: str) -> str:
  11. """
  12. 企业级标准文本净化器 (Text Sanitizer)。
  13. 1. 规范化 (Normalization):统一转换为标准 NFC 格式。
  14. 2. 剔除/替换不合法及不可见字节:过滤 NUL、零宽字符、BOM 标签等。
  15. 3. 清洗非打印字符及乱码:剔除 C0/C1 控制符号,剥离 Unicode PUA 私有使用区乱码符号。
  16. 4. 编码收敛防爆:利用 utf-8 ignore 安全剥离孤立的 UTF-16 代理项。
  17. """
  18. if not text:
  19. return ""
  20. text = unicodedata.normalize("NFC", text)
  21. text = _INVISIBLE_CHAR_RE.sub("", text)
  22. text = _CONTROL_CHAR_RE.sub("", text)
  23. text = re.sub(r"[\ue000-\uf8ff]", "", text)
  24. try:
  25. cleaned = text.encode("utf-8", "ignore").decode("utf-8", "ignore")
  26. except Exception:
  27. chars = []
  28. for char in text:
  29. if 0xD800 <= ord(char) <= 0xDFFF:
  30. continue
  31. chars.append(char)
  32. cleaned = "".join(chars)
  33. return cleaned
  34. # 分隔符按语义粒度分组,越靠前越优先作为分块边界。
  35. # 切分时保留分隔符作为独立原子,保证拼接后原文不变。
  36. _SENTENCE_SEPARATORS = ["\n\n", "。", "!", "?", ";", "\n"]
  37. _CLAUSE_SEPARATORS = [",", "、"]
  38. _WORD_SEPARATORS = [" ", ""]
  39. DEFAULT_SEPARATORS = _SENTENCE_SEPARATORS + _CLAUSE_SEPARATORS + _WORD_SEPARATORS
  40. class SemanticTextSplitter:
  41. """
  42. 语义优先的文本分块器。
  43. 核心保证:
  44. 1. 分块后所有 chunk 按顺序拼接,等于原始输入。
  45. 2. 每个 chunk 的末尾尽量落在语义边界(段落 > 句子 > 子句 > 词)。
  46. 3. 重叠区只在边界处截取,不会从句子中间截断。
  47. """
  48. def __init__(
  49. self,
  50. chunk_size: int = 800,
  51. chunk_overlap: int = 100,
  52. separators: List[str] | None = None,
  53. use_semantic_chunking: bool = False,
  54. semantic_similarity_threshold: float = 0.6,
  55. ):
  56. self.chunk_size = chunk_size
  57. self.chunk_overlap = chunk_overlap
  58. self.separators = separators or DEFAULT_SEPARATORS
  59. self.use_semantic_chunking = use_semantic_chunking
  60. self.semantic_chunker = None
  61. if use_semantic_chunking:
  62. # overlap 在语义段落边界处由 SemanticChunker 处理,保证承上启下句子完整
  63. self.semantic_chunker = SemanticChunker(
  64. similarity_threshold=semantic_similarity_threshold,
  65. max_chunk_size=chunk_size,
  66. chunk_overlap=chunk_overlap,
  67. )
  68. def split_text(self, text: str) -> List[str]:
  69. """对单段文本执行分块。"""
  70. text = sanitize_text(text)
  71. if not text:
  72. return []
  73. if len(text) <= self.chunk_size:
  74. return [text]
  75. # 若启用语义分块,先按话题边界粗分,再在每个语义段落内做规则细分。
  76. # 保证所有段落拼接仍等于原文。
  77. if self.use_semantic_chunking and self.semantic_chunker is not None:
  78. semantic_chunks = self.semantic_chunker.split_text(text)
  79. result: List[str] = []
  80. for chunk in semantic_chunks:
  81. if len(chunk) <= self.chunk_size:
  82. result.append(chunk)
  83. else:
  84. atoms = self._split_into_atoms(chunk)
  85. result.extend(self._merge_atoms(atoms))
  86. return result
  87. # 未启用语义分块:按标点规则切分。
  88. atoms = self._split_into_atoms(text)
  89. return self._merge_atoms(atoms)
  90. def _split_into_atoms(self, text: str) -> List[str]:
  91. """
  92. 递归地在最粗可用的分隔符处切分文本,保留分隔符作为独立原子。
  93. 例如 "。" 切分后,标点本身会作为一个原子项,保证后续拼接不丢标点。
  94. """
  95. return self._split_recursive(text, self.separators.copy())
  96. def _split_recursive(self, text: str, separators: List[str]) -> List[str]:
  97. if len(text) <= self.chunk_size or not separators:
  98. return [text] if text else []
  99. separator = separators.pop(0)
  100. if separator == "":
  101. # 兜底:按字符切分
  102. return list(text)
  103. parts = text.split(separator)
  104. if len(parts) <= 1:
  105. return self._split_recursive(text, separators)
  106. result: List[str] = []
  107. for i, part in enumerate(parts):
  108. if i > 0:
  109. # 分隔符本身作为独立原子保留
  110. result.append(separator)
  111. if not part:
  112. continue
  113. if len(part) > self.chunk_size:
  114. result.extend(self._split_recursive(part, separators.copy()))
  115. else:
  116. result.append(part)
  117. return result
  118. def _merge_atoms(self, atoms: List[str]) -> List[str]:
  119. """
  120. 把原子片段合并成接近 chunk_size 的块。
  121. 合并规则:
  122. 1. 原子不可拆分;
  123. 2. 中间 chunk 的末尾尽量落在语义分隔符上;
  124. 3. 因回退而移出当前 chunk 的非边界原子,会保留到下一个 chunk 中,保证原文不丢失;
  125. 4. 若启用 overlap,下一个 chunk 以当前 chunk 末尾的边界感知重叠区开头。
  126. """
  127. if not atoms:
  128. return []
  129. chunks: List[str] = []
  130. current: List[str] = []
  131. current_len = 0
  132. def last_sep_index(seq: List[str]) -> int:
  133. for idx in range(len(seq) - 1, -1, -1):
  134. if seq[idx] in self.separators:
  135. return idx
  136. return -1
  137. for atom in atoms:
  138. atom_len = len(atom)
  139. # 加入当前原子会超过 chunk_size,则先结束当前 chunk
  140. if current and current_len + atom_len > self.chunk_size:
  141. sep_idx = last_sep_index(current)
  142. if sep_idx >= 0:
  143. chunk_atoms = current[: sep_idx + 1]
  144. pending = current[sep_idx + 1 :] + [atom]
  145. else:
  146. # 当前窗口内没有可用分隔符,直接结束整个窗口
  147. chunk_atoms = current
  148. pending = [atom]
  149. chunks.append("".join(chunk_atoms))
  150. if self.chunk_overlap > 0 and chunk_atoms:
  151. overlap_atoms = self._extract_overlap_atoms(chunk_atoms)
  152. current = overlap_atoms + pending
  153. else:
  154. current = pending
  155. current_len = sum(len(a) for a in current)
  156. continue
  157. current.append(atom)
  158. current_len += atom_len
  159. if current:
  160. chunks.append("".join(current))
  161. return chunks
  162. def _extract_overlap_atoms(self, atoms: List[str]) -> List[str]:
  163. """
  164. 从一组原子末尾回退,截取不超过 chunk_overlap 且以语义边界结尾的重叠区。
  165. 优先让整个重叠区以句子或子句分隔符结尾。
  166. """
  167. if self.chunk_overlap <= 0 or not atoms:
  168. return []
  169. overlap: List[str] = []
  170. length = 0
  171. for atom in reversed(atoms):
  172. overlap.insert(0, atom)
  173. length += len(atom)
  174. if length >= self.chunk_overlap:
  175. break
  176. # 让 overlap 以语义边界结束
  177. while len(overlap) > 1:
  178. if overlap[-1] in self.separators:
  179. break
  180. removed_len = len(overlap[-1])
  181. if length - removed_len >= self.chunk_overlap // 3:
  182. length -= removed_len
  183. overlap.pop()
  184. else:
  185. break
  186. return overlap
  187. @staticmethod
  188. def _atoms_length(atoms: List[str]) -> int:
  189. return sum(len(a) for a in atoms)
  190. class HierarchicalTextSplitter:
  191. """三级层次化分块(L1 / L2 / L3)"""
  192. def __init__(self, chunk_size: int = 800, chunk_overlap: int = 100):
  193. # L1 是对原文的粗粒度切分,L1 之间不允许 overlap,保证所有 L1 拼接等于原文。
  194. # L2/L3 在父级内部进行切分,启用基于 embedding 相似度的语义边界;
  195. # 允许 L2/L3 在语义边界处保留 overlap,使承上启下的句子同时出现在前后块中。
  196. level_1_size = max(2000, chunk_size * 3)
  197. level_2_size = max(1000, chunk_size * 2)
  198. level_3_size = max(600, chunk_size)
  199. self._splitter_level_1 = SemanticTextSplitter(
  200. chunk_size=level_1_size,
  201. chunk_overlap=0,
  202. separators=DEFAULT_SEPARATORS,
  203. use_semantic_chunking=False,
  204. )
  205. self._splitter_level_2 = SemanticTextSplitter(
  206. chunk_size=level_2_size,
  207. chunk_overlap=chunk_overlap,
  208. separators=DEFAULT_SEPARATORS,
  209. use_semantic_chunking=True,
  210. )
  211. self._splitter_level_3 = SemanticTextSplitter(
  212. chunk_size=level_3_size,
  213. chunk_overlap=chunk_overlap,
  214. separators=DEFAULT_SEPARATORS,
  215. use_semantic_chunking=True,
  216. )
  217. @staticmethod
  218. def _build_chunk_id(document_id: int, level: int, index: int) -> str:
  219. return f"{document_id}::l{level}::{index}"
  220. def split_text(
  221. self,
  222. text: str,
  223. document_id: int,
  224. filename: str = "",
  225. file_type: str = "",
  226. file_path: str = "",
  227. page_number: int = 0,
  228. ) -> List[Dict]:
  229. """对单段文本执行三级分块,返回 L1/L2/L3 全部分块。"""
  230. text = sanitize_text(text)
  231. if not text:
  232. return []
  233. base_doc = {
  234. "filename": sanitize_text(filename),
  235. "file_path": sanitize_text(file_path),
  236. "file_type": sanitize_text(file_type),
  237. "page_number": int(page_number),
  238. }
  239. root_chunks: List[Dict] = []
  240. level_1_counter = 0
  241. level_2_counter = 0
  242. level_3_counter = 0
  243. chunk_idx = 0
  244. level_1_texts = self._splitter_level_1.split_text(text)
  245. for level_1_text in level_1_texts:
  246. if not level_1_text:
  247. continue
  248. level_1_id = self._build_chunk_id(document_id, 1, level_1_counter)
  249. level_1_counter += 1
  250. level_1_chunk = {
  251. **base_doc,
  252. "text": level_1_text,
  253. "chunk_id": level_1_id,
  254. "parent_chunk_id": "",
  255. "root_chunk_id": level_1_id,
  256. "chunk_level": 1,
  257. "chunk_idx": chunk_idx,
  258. }
  259. chunk_idx += 1
  260. root_chunks.append(level_1_chunk)
  261. level_2_texts = self._splitter_level_2.split_text(level_1_text)
  262. for level_2_text in level_2_texts:
  263. if not level_2_text:
  264. continue
  265. level_2_id = self._build_chunk_id(document_id, 2, level_2_counter)
  266. level_2_counter += 1
  267. level_2_chunk = {
  268. **base_doc,
  269. "text": level_2_text,
  270. "chunk_id": level_2_id,
  271. "parent_chunk_id": level_1_id,
  272. "root_chunk_id": level_1_id,
  273. "chunk_level": 2,
  274. "chunk_idx": chunk_idx,
  275. }
  276. chunk_idx += 1
  277. root_chunks.append(level_2_chunk)
  278. level_3_texts = self._splitter_level_3.split_text(level_2_text)
  279. for level_3_text in level_3_texts:
  280. if not level_3_text:
  281. continue
  282. level_3_id = self._build_chunk_id(document_id, 3, level_3_counter)
  283. level_3_counter += 1
  284. root_chunks.append({
  285. **base_doc,
  286. "text": level_3_text,
  287. "chunk_id": level_3_id,
  288. "parent_chunk_id": level_2_id,
  289. "root_chunk_id": level_1_id,
  290. "chunk_level": 3,
  291. "chunk_idx": chunk_idx,
  292. })
  293. chunk_idx += 1
  294. return root_chunks