Files
RAG-CUT/backend/rag_cut/splitters/parent_child.py
T

104 lines
3.8 KiB
Python
Raw Normal View History

2026-07-16 11:12:17 +08:00
"""Parent/child delimiter splitting for fine retrieval + coarse recall."""
from __future__ import annotations
from rag_cut.models import Block, SplitConfig
from rag_cut.splitters.default_splitter import _split_by_size
from rag_cut.splitters.delimiter import partition_blocks_by_delimiter
from rag_cut.splitters.heading_splitter import ChunkGroup
CHUNK_STRATEGY = "parent_child_delimiter"
CHILD_MAX_HARD_LIMIT = 1500
def _rendered_len(blocks: list[Block]) -> int:
return sum(len(b.render()) + 2 for b in blocks)
def _validate(config: SplitConfig) -> tuple[str, str | None, int, int]:
parent_delimiter = (config.parent_delimiter or config.delimiter or "").strip()
if not parent_delimiter:
raise ValueError("parent_delimiter is required for parent_child split mode")
child_delimiter = (config.child_delimiter or "").strip() or None
parent_max = max(200, int(config.max_chunk_size or 1500))
child_max = int(config.child_max_size or 512)
child_max = max(50, min(child_max, CHILD_MAX_HARD_LIMIT, parent_max))
return parent_delimiter, child_delimiter, parent_max, child_max
def _size_cap(groups: list[list[Block]], max_size: int, overlap: int) -> list[list[Block]]:
sized: list[list[Block]] = []
size_config = SplitConfig(max_chunk_size=max_size, overlap=overlap)
for group in groups:
if _rendered_len(group) <= max_size:
sized.append(group)
else:
sized.extend(_split_by_size(group, size_config))
return sized or groups
def split_by_parent_child(blocks: list[Block], config: SplitConfig) -> list[ChunkGroup]:
"""
Split into parent chunks (context) and child chunks (retrieval).
1. Partition by parent_delimiter (then cap by parent max length).
2. For each parent: keep a full parent chunk (retrieval=false).
3. Partition parent by child_delimiter (or by length) into children
capped by child_max_size (retrieval=true, linked via parent_section_id).
"""
if not blocks:
return []
parent_delimiter, child_delimiter, parent_max, child_max = _validate(config)
overlap = max(0, int(config.overlap or 0))
parents = partition_blocks_by_delimiter(blocks, parent_delimiter)
if not parents:
parents = [blocks]
parents = _size_cap(parents, parent_max, overlap)
groups: list[ChunkGroup] = []
for parent_idx, parent_blocks in enumerate(parents):
section_id = f"pc-{parent_idx}"
groups.append(
ChunkGroup(
blocks=list(parent_blocks),
meta={
"section_id": section_id,
"is_section_parent": True,
"is_sub_chunk": False,
"retrieval": False,
"chunk_strategy": f"{CHUNK_STRATEGY}_parent",
"parent_index": parent_idx,
},
)
)
if child_delimiter:
children = partition_blocks_by_delimiter(parent_blocks, child_delimiter)
else:
children = [parent_blocks]
if not children:
children = [parent_blocks]
children = _size_cap(children, child_max, overlap)
for child_idx, child_blocks in enumerate(children):
groups.append(
ChunkGroup(
blocks=list(child_blocks),
meta={
"section_id": f"{section_id}-sub-{child_idx}",
"parent_section_id": section_id,
"is_section_parent": False,
"is_sub_chunk": True,
"retrieval": True,
"chunk_strategy": f"{CHUNK_STRATEGY}_child",
"parent_index": parent_idx,
"sub_chunk_index": child_idx,
},
)
)
return groups