优化切片,增加父节点章节、图片描述
This commit is contained in:
parent
626139226d
commit
129846dc17
468
chunk_text.py
468
chunk_text.py
@ -24,6 +24,7 @@ import threading
|
|||||||
import aiofiles
|
import aiofiles
|
||||||
import httpx
|
import httpx
|
||||||
import logging
|
import logging
|
||||||
|
from resetlevel import reset_textlevel
|
||||||
from rapidocr_onnxruntime import RapidOCR
|
from rapidocr_onnxruntime import RapidOCR
|
||||||
from PIL import Image, ImageOps, UnidentifiedImageError
|
from PIL import Image, ImageOps, UnidentifiedImageError
|
||||||
from config import (
|
from config import (
|
||||||
@ -51,6 +52,8 @@ _image_ocr_cache: Dict[str, str] = {}
|
|||||||
_image_ocr_cache_lock = threading.Lock()
|
_image_ocr_cache_lock = threading.Lock()
|
||||||
_api_url_inflight = [0] * len(API_URLS)
|
_api_url_inflight = [0] * len(API_URLS)
|
||||||
_api_url_lock = threading.Lock()
|
_api_url_lock = threading.Lock()
|
||||||
|
IMAGE_OCR_PREFIX = "图片中文本内容为:"
|
||||||
|
IMAGE_OCR_PREFIXES = (IMAGE_OCR_PREFIX, "图片包含的文字内容为:")
|
||||||
|
|
||||||
|
|
||||||
def _image_mime_type(suffix: str) -> str:
|
def _image_mime_type(suffix: str) -> str:
|
||||||
@ -250,7 +253,7 @@ def extract_image_text_sync(image_path: Union[str, PathLib]) -> dict:
|
|||||||
ocr_result, elapse = ocr_engine(resolved_path)
|
ocr_result, elapse = ocr_engine(resolved_path)
|
||||||
|
|
||||||
texts = [item[1] for item in (ocr_result or [])]
|
texts = [item[1] for item in (ocr_result or [])]
|
||||||
full_text = "图片包含的文字内容为:" + ",".join(texts) if texts else ""
|
full_text = IMAGE_OCR_PREFIX + ",".join(texts) if texts else ""
|
||||||
with _image_ocr_cache_lock:
|
with _image_ocr_cache_lock:
|
||||||
_image_ocr_cache[resolved_path] = full_text
|
_image_ocr_cache[resolved_path] = full_text
|
||||||
return {
|
return {
|
||||||
@ -261,7 +264,7 @@ def extract_image_text_sync(image_path: Union[str, PathLib]) -> dict:
|
|||||||
|
|
||||||
|
|
||||||
def build_image_markdown_with_ocr(markdown_image: str, local_image_path: Optional[PathLib]) -> str:
|
def build_image_markdown_with_ocr(markdown_image: str, local_image_path: Optional[PathLib]) -> str:
|
||||||
if "图片包含的文字内容为:" in markdown_image:
|
if any(prefix in markdown_image for prefix in IMAGE_OCR_PREFIXES):
|
||||||
return markdown_image
|
return markdown_image
|
||||||
if not local_image_path:
|
if not local_image_path:
|
||||||
return markdown_image
|
return markdown_image
|
||||||
@ -278,6 +281,179 @@ def build_image_markdown_with_ocr(markdown_image: str, local_image_path: Optiona
|
|||||||
return f"{markdown_image}\n{ocr_full_text}"
|
return f"{markdown_image}\n{ocr_full_text}"
|
||||||
|
|
||||||
|
|
||||||
|
async def get_image_ocr_text(local_image_path: Optional[PathLib]) -> str:
|
||||||
|
if not local_image_path:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
try:
|
||||||
|
ocr_data = await extract_image_text(local_image_path)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(f"Failed to OCR local image {local_image_path}: {exc}")
|
||||||
|
return ""
|
||||||
|
|
||||||
|
return (ocr_data.get("full_text") or "").strip()
|
||||||
|
|
||||||
|
|
||||||
|
def stringify_content_field(value: Any) -> str:
|
||||||
|
if value is None:
|
||||||
|
return ""
|
||||||
|
if isinstance(value, list):
|
||||||
|
return "\n".join(str(item) for item in value if str(item).strip())
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return json.dumps(value, ensure_ascii=False)
|
||||||
|
return str(value)
|
||||||
|
|
||||||
|
|
||||||
|
def clean_table_body(table_body: Any) -> str:
|
||||||
|
table_body = stringify_content_field(table_body)
|
||||||
|
for pattern in [' colspan="1"', ' rowspan="1"', " colspan='1'", " rowspan='1'", " colspan=1", " rowspan=1"]:
|
||||||
|
table_body = table_body.replace(pattern, "")
|
||||||
|
return table_body
|
||||||
|
|
||||||
|
|
||||||
|
def join_nonempty_parts(*parts: Any) -> str:
|
||||||
|
return "\n".join(part for part in (stringify_content_field(part).strip() for part in parts) if part)
|
||||||
|
|
||||||
|
|
||||||
|
def record_to_chunk_text(ins: Dict[str, Any]) -> str:
|
||||||
|
record_type = ins.get("type")
|
||||||
|
ocr_text = ins.get("ocr_text", "")
|
||||||
|
|
||||||
|
if record_type == "text":
|
||||||
|
text = stringify_content_field(ins.get("text", ""))
|
||||||
|
level = ins.get("text_level")
|
||||||
|
if level is not None and isinstance(level, int) and level > 0:
|
||||||
|
text = "#" * min(level, 6) + " " + text
|
||||||
|
return join_nonempty_parts(ocr_text, text)
|
||||||
|
|
||||||
|
if record_type == "table":
|
||||||
|
return join_nonempty_parts(
|
||||||
|
ocr_text,
|
||||||
|
ins.get("table_caption"),
|
||||||
|
clean_table_body(ins.get("table_body", "")),
|
||||||
|
ins.get("table_footnote"),
|
||||||
|
)
|
||||||
|
|
||||||
|
if record_type == "image":
|
||||||
|
return join_nonempty_parts(
|
||||||
|
ocr_text,
|
||||||
|
ins.get("image_caption"),
|
||||||
|
ins.get("img_path"),
|
||||||
|
ins.get("image_footnote"),
|
||||||
|
)
|
||||||
|
|
||||||
|
if record_type == "chart":
|
||||||
|
return join_nonempty_parts(
|
||||||
|
ocr_text,
|
||||||
|
ins.get("chart_caption"),
|
||||||
|
ins.get("content"),
|
||||||
|
ins.get("img_path"),
|
||||||
|
ins.get("chart_footnote"),
|
||||||
|
)
|
||||||
|
|
||||||
|
if record_type == "equation":
|
||||||
|
return join_nonempty_parts(
|
||||||
|
ocr_text,
|
||||||
|
ins.get("text"),
|
||||||
|
ins.get("img_path"),
|
||||||
|
)
|
||||||
|
|
||||||
|
if record_type == "code":
|
||||||
|
return join_nonempty_parts(
|
||||||
|
ocr_text,
|
||||||
|
ins.get("code_caption"),
|
||||||
|
ins.get("code_body"),
|
||||||
|
ins.get("code_footnote"),
|
||||||
|
)
|
||||||
|
|
||||||
|
if record_type == "list":
|
||||||
|
return join_nonempty_parts(ocr_text, ins.get("list_items"))
|
||||||
|
|
||||||
|
if record_type in {"discarded", "header", "footer", "page_number"}:
|
||||||
|
return join_nonempty_parts(ocr_text, ins.get("text"))
|
||||||
|
|
||||||
|
return join_nonempty_parts(ocr_text, ins.get("text"))
|
||||||
|
|
||||||
|
|
||||||
|
def is_supported_chunk_record(ins: Dict[str, Any]) -> bool:
|
||||||
|
return ins.get("type") in {
|
||||||
|
"text",
|
||||||
|
"table",
|
||||||
|
"image",
|
||||||
|
"chart",
|
||||||
|
"equation",
|
||||||
|
"code",
|
||||||
|
"list",
|
||||||
|
"discarded",
|
||||||
|
"header",
|
||||||
|
"footer",
|
||||||
|
"page_number",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def append_caption_ocr(ins: Dict[str, Any], ocr_texts: List[str]) -> None:
|
||||||
|
caption_field = {
|
||||||
|
"image": "image_caption",
|
||||||
|
"table": "table_caption",
|
||||||
|
"chart": "chart_caption",
|
||||||
|
"code": "code_caption",
|
||||||
|
}.get(ins.get("type"))
|
||||||
|
if not caption_field:
|
||||||
|
caption_field = "ocr_text"
|
||||||
|
|
||||||
|
caption = ins.get(caption_field, [])
|
||||||
|
existing_text = stringify_content_field(caption)
|
||||||
|
if caption_field == "ocr_text":
|
||||||
|
existing_text = join_nonempty_parts(existing_text, ins.get("text"))
|
||||||
|
|
||||||
|
additions = []
|
||||||
|
for text in ocr_texts:
|
||||||
|
text = (text or "").strip()
|
||||||
|
text_without_prefix = text
|
||||||
|
for prefix in IMAGE_OCR_PREFIXES:
|
||||||
|
if text_without_prefix.startswith(prefix):
|
||||||
|
text_without_prefix = text_without_prefix[len(prefix):].strip()
|
||||||
|
break
|
||||||
|
|
||||||
|
if (
|
||||||
|
text
|
||||||
|
and text not in existing_text
|
||||||
|
and text_without_prefix not in existing_text
|
||||||
|
and text not in additions
|
||||||
|
):
|
||||||
|
additions.append(text)
|
||||||
|
|
||||||
|
if not additions:
|
||||||
|
return
|
||||||
|
|
||||||
|
if caption_field == "ocr_text":
|
||||||
|
ins[caption_field] = join_nonempty_parts(additions, caption)
|
||||||
|
return
|
||||||
|
|
||||||
|
if isinstance(caption, list):
|
||||||
|
ins[caption_field] = additions + caption
|
||||||
|
elif caption is None or not str(caption).strip():
|
||||||
|
ins[caption_field] = additions
|
||||||
|
else:
|
||||||
|
ins[caption_field] = "\n".join(additions) + "\n" + str(caption).lstrip()
|
||||||
|
|
||||||
|
|
||||||
|
def build_image_markdown(markdown_or_path: str, prefix: str) -> Tuple[str, Optional[PathLib]]:
|
||||||
|
img_path = str(markdown_or_path).strip()
|
||||||
|
if not img_path:
|
||||||
|
return "", None
|
||||||
|
|
||||||
|
if img_path.strip().startswith("!["):
|
||||||
|
return img_path, resolve_local_image_path(img_path)
|
||||||
|
|
||||||
|
p = PurePosixPath(img_path)
|
||||||
|
if p.parts and p.parts[0] == "images":
|
||||||
|
relative_path = str(PurePosixPath(*p.parts[1:]))
|
||||||
|
return f"", resolve_local_image_path(relative_path)
|
||||||
|
|
||||||
|
return f"", resolve_local_image_path(img_path)
|
||||||
|
|
||||||
|
|
||||||
def image_base64_to_data_url(image_base64: Union[str, bytes], suffix: str) -> str:
|
def image_base64_to_data_url(image_base64: Union[str, bytes], suffix: str) -> str:
|
||||||
if isinstance(image_base64, bytes):
|
if isinstance(image_base64, bytes):
|
||||||
raw = base64.b64encode(image_base64).decode("ascii")
|
raw = base64.b64encode(image_base64).decode("ascii")
|
||||||
@ -602,76 +778,14 @@ def split_content_bbox(data, max_length=8000):
|
|||||||
bbox = ins.get("bbox")
|
bbox = ins.get("bbox")
|
||||||
page_idx = ins.get("page_idx", -1)
|
page_idx = ins.get("page_idx", -1)
|
||||||
|
|
||||||
new_text = ""
|
if not is_supported_chunk_record(ins):
|
||||||
should_record_highlight = False # 标记是否需要记录 highlight(即使 new_text 为空)
|
continue
|
||||||
|
|
||||||
if ins["type"] == "text":
|
new_text = record_to_chunk_text(ins)
|
||||||
text = ins.get("text", "")
|
if new_text.strip():
|
||||||
level = ins.get("text_level")
|
new_text += "\n"
|
||||||
if level is not None and isinstance(level, int) and level > 0:
|
|
||||||
text = "#" * min(level, 6) + " " + text # 安全限制标题级别
|
|
||||||
new_text = text + "\n"
|
|
||||||
should_record_highlight = True
|
should_record_highlight = True
|
||||||
|
|
||||||
elif ins["type"] == "table":
|
|
||||||
table_caption = ins.get("table_caption")
|
|
||||||
table_body = ins.get("table_body", "")
|
|
||||||
|
|
||||||
# 清理冗余属性
|
|
||||||
for pattern in [' colspan="1"', ' rowspan="1"', " colspan='1'", " rowspan='1'", " colspan=1", " rowspan=1"]:
|
|
||||||
table_body = table_body.replace(pattern, "")
|
|
||||||
|
|
||||||
caption_str = ""
|
|
||||||
if table_caption:
|
|
||||||
if isinstance(table_caption, list):
|
|
||||||
caption_str = "\n".join(str(c) for c in table_caption)
|
|
||||||
else:
|
|
||||||
caption_str = str(table_caption)
|
|
||||||
|
|
||||||
# 判断是否有实质内容
|
|
||||||
has_caption = bool(caption_str.strip())
|
|
||||||
has_body = bool(table_body.strip())
|
|
||||||
|
|
||||||
if has_caption or has_body:
|
|
||||||
new_text = caption_str + ("\n" if has_caption else "") + table_body + "\n"
|
|
||||||
else:
|
|
||||||
new_text = "" # 空内容,不拼接
|
|
||||||
|
|
||||||
# 即使内容为空,只要是 table 类型,就记录 highlight(满足你的需求)
|
|
||||||
should_record_highlight = True
|
|
||||||
|
|
||||||
elif ins["type"] == "image":
|
|
||||||
image_caption = ins.get("image_caption", "")
|
|
||||||
img_path = ins.get("img_path", "")
|
|
||||||
|
|
||||||
if isinstance(image_caption, list):
|
|
||||||
image_caption = "\n".join(str(item) for item in image_caption)
|
|
||||||
elif not isinstance(image_caption, str):
|
|
||||||
image_caption = str(image_caption)
|
|
||||||
|
|
||||||
if isinstance(img_path, list):
|
|
||||||
img_path = "\n".join(str(p) for p in img_path)
|
|
||||||
elif not isinstance(img_path, str):
|
|
||||||
img_path = str(img_path)
|
|
||||||
|
|
||||||
if image_caption.strip() or img_path.strip():
|
|
||||||
new_text = image_caption + "\n" + img_path + "\n"
|
|
||||||
else:
|
|
||||||
new_text = ""
|
|
||||||
|
|
||||||
should_record_highlight = True
|
|
||||||
|
|
||||||
elif ins["type"] == "discarded":
|
|
||||||
text = ins.get("text", "")
|
|
||||||
if text.strip():
|
|
||||||
new_text = text + "\n"
|
|
||||||
else:
|
|
||||||
new_text = ""
|
|
||||||
should_record_highlight = True
|
|
||||||
|
|
||||||
else:
|
|
||||||
continue # 不支持的 type,跳过且不记录 highlight
|
|
||||||
|
|
||||||
# --- 关键:即使 new_text 为空,只要 should_record_highlight 为 True,就要处理 chunk 切分和 highlight 记录 ---
|
# --- 关键:即使 new_text 为空,只要 should_record_highlight 为 True,就要处理 chunk 切分和 highlight 记录 ---
|
||||||
|
|
||||||
# 构造 highlight 项(即使 new_text 为空)
|
# 构造 highlight 项(即使 new_text 为空)
|
||||||
@ -707,7 +821,7 @@ def split_content_bbox(data, max_length=8000):
|
|||||||
|
|
||||||
|
|
||||||
|
|
||||||
def group_records(records):
|
def _group_records_without_parent_context(records):
|
||||||
result = []
|
result = []
|
||||||
current_group = []
|
current_group = []
|
||||||
|
|
||||||
@ -728,10 +842,50 @@ def group_records(records):
|
|||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _get_group_level(group):
|
||||||
|
if not group:
|
||||||
|
return None
|
||||||
|
|
||||||
|
level = group[0].get("text_level")
|
||||||
|
if isinstance(level, int) and level > 0:
|
||||||
|
return level
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def group_records(records):
|
||||||
|
"""
|
||||||
|
Split records by text_level, then prepend ancestor section content.
|
||||||
|
|
||||||
|
For example, a text_level=2 section will include the nearest text_level=1
|
||||||
|
group content before its own content. A text_level=3 section will include
|
||||||
|
the nearest level 1 and level 2 group content before its own content.
|
||||||
|
"""
|
||||||
|
base_groups = _group_records_without_parent_context(records)
|
||||||
|
result = []
|
||||||
|
section_stack = [] # [(level, merged_group)]
|
||||||
|
|
||||||
|
for group in base_groups:
|
||||||
|
level = _get_group_level(group)
|
||||||
|
if level is None:
|
||||||
|
result.append(group)
|
||||||
|
continue
|
||||||
|
|
||||||
|
while section_stack and section_stack[-1][0] >= level:
|
||||||
|
section_stack.pop()
|
||||||
|
|
||||||
|
parent_group = section_stack[-1][1] if section_stack else []
|
||||||
|
merged_group = parent_group + group
|
||||||
|
result.append(merged_group)
|
||||||
|
section_stack.append((level, merged_group))
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
def get_chunk_bbox_otherfile(data, max_length=8000):
|
def get_chunk_bbox_otherfile(data, max_length=8000):
|
||||||
"""
|
"""
|
||||||
处理没有 bbox 信息的文件数据(主要面向纯表格类内容)。
|
处理没有 bbox 信息的文件数据。
|
||||||
每个 table 作为一个独立切片,不与其他记录拼接。
|
每条支持的记录作为一个独立切片,不与其他记录拼接。
|
||||||
|
|
||||||
- 单个 table 内容超过 max_length 时,调用 chunk_reset 在 </tr> 处安全切分,
|
- 单个 table 内容超过 max_length 时,调用 chunk_reset 在 </tr> 处安全切分,
|
||||||
切分后的多个 chunk 共享同一组 positions
|
切分后的多个 chunk 共享同一组 positions
|
||||||
@ -743,38 +897,14 @@ def get_chunk_bbox_otherfile(data, max_length=8000):
|
|||||||
slices = []
|
slices = []
|
||||||
|
|
||||||
for ins in data:
|
for ins in data:
|
||||||
if ins.get("type") != "table":
|
if not is_supported_chunk_record(ins):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
page_idx = ins.get("page_idx", -1)
|
page_idx = ins.get("page_idx", -1)
|
||||||
bbox = ins.get("bbox", [])
|
bbox = ins.get("bbox", [])
|
||||||
|
content = record_to_chunk_text(ins)
|
||||||
table_caption = ins.get("table_caption")
|
if not content.strip():
|
||||||
table_body = ins.get("table_body", "") or ""
|
continue
|
||||||
|
|
||||||
# 清理冗余属性
|
|
||||||
for pattern in [
|
|
||||||
' colspan="1"', ' rowspan="1"',
|
|
||||||
" colspan='1'", " rowspan='1'",
|
|
||||||
" colspan=1", " rowspan=1",
|
|
||||||
]:
|
|
||||||
table_body = table_body.replace(pattern, "")
|
|
||||||
|
|
||||||
# 拼接 caption
|
|
||||||
caption_str = ""
|
|
||||||
if table_caption:
|
|
||||||
if isinstance(table_caption, list):
|
|
||||||
caption_str = "\n".join(str(c) for c in table_caption)
|
|
||||||
else:
|
|
||||||
caption_str = str(table_caption)
|
|
||||||
|
|
||||||
has_caption = bool(caption_str.strip())
|
|
||||||
has_body = bool(table_body.strip())
|
|
||||||
|
|
||||||
if not (has_caption or has_body):
|
|
||||||
continue # 空表格跳过
|
|
||||||
|
|
||||||
content = caption_str + ("\n" if has_caption else "") + table_body
|
|
||||||
positions = [{"page_idx": page_idx, "bbox": bbox}]
|
positions = [{"page_idx": page_idx, "bbox": bbox}]
|
||||||
|
|
||||||
# 单表超长 → 在 </tr> 处切分,多个 chunk 共享 positions
|
# 单表超长 → 在 </tr> 处切分,多个 chunk 共享 positions
|
||||||
@ -875,38 +1005,11 @@ def get_chunk_bbox(data):
|
|||||||
bbox = ins["bbox"]
|
bbox = ins["bbox"]
|
||||||
positions = [{"page_idx": page_idx, "bbox": bbox}]
|
positions = [{"page_idx": page_idx, "bbox": bbox}]
|
||||||
|
|
||||||
if ins["type"] == "text":
|
if not is_supported_chunk_record(ins):
|
||||||
if ins.get("text_level"):
|
continue
|
||||||
text = ins.get("text", "")
|
|
||||||
level = ins.get("text_level")
|
|
||||||
content = "#" * level + text
|
|
||||||
else:
|
|
||||||
content = ins.get("text", "")
|
|
||||||
elif ins["type"] == "table":
|
|
||||||
content = (ins.get("table_caption") or "") + "\n" + (ins.get("table_body") or "")
|
|
||||||
elif ins["type"] == "image":
|
|
||||||
image_caption = ins.get("image_caption", [])
|
|
||||||
img_path = ins.get("img_path", [])
|
|
||||||
|
|
||||||
# 如果 image_caption 是列表,则将其转换为字符串
|
content = record_to_chunk_text(ins)
|
||||||
if isinstance(image_caption, list):
|
if not content.strip():
|
||||||
image_caption = "\n".join(str(item) for item in image_caption)
|
|
||||||
elif not isinstance(image_caption, str):
|
|
||||||
image_caption = str(image_caption)
|
|
||||||
|
|
||||||
# 如果 img_path 是列表,则将其转换为字符串
|
|
||||||
if isinstance(img_path, list):
|
|
||||||
img_path = "\n".join(str(p) for p in img_path)
|
|
||||||
elif not isinstance(img_path, str):
|
|
||||||
img_path = str(img_path)
|
|
||||||
|
|
||||||
if image_caption.strip() or img_path.strip():
|
|
||||||
content = (image_caption or "") + "\n" + (img_path or "")
|
|
||||||
else:
|
|
||||||
continue # 跳过空图像项
|
|
||||||
elif ins["type"] == "discarded":
|
|
||||||
content = ins.get("text", "")
|
|
||||||
else:
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
slices.append({"content": content, "positions": positions})
|
slices.append({"content": content, "positions": positions})
|
||||||
@ -977,7 +1080,7 @@ def split_text_preserve_sentences(text, max_length=7500):
|
|||||||
|
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
def data_replace(data, prefix):
|
async def data_replace(data, prefix):
|
||||||
"""
|
"""
|
||||||
遍历 data 列表,若元素包含 'img_path' 字段,则移除前缀 'images/' 并拼接新前缀。
|
遍历 data 列表,若元素包含 'img_path' 字段,则移除前缀 'images/' 并拼接新前缀。
|
||||||
使用 pathlib 安全处理路径。
|
使用 pathlib 安全处理路径。
|
||||||
@ -985,32 +1088,27 @@ def data_replace(data, prefix):
|
|||||||
for ins in data:
|
for ins in data:
|
||||||
if isinstance(ins, dict) and "img_path" in ins:
|
if isinstance(ins, dict) and "img_path" in ins:
|
||||||
img_path = ins["img_path"]
|
img_path = ins["img_path"]
|
||||||
|
local_image_paths = []
|
||||||
if isinstance(img_path, list):
|
if isinstance(img_path, list):
|
||||||
ins["img_path"] = "\n".join(
|
markdown_images = []
|
||||||
data_replace([{"img_path": str(item)}], prefix)[0]["img_path"]
|
for item in img_path:
|
||||||
for item in img_path
|
if not str(item).strip():
|
||||||
if str(item).strip()
|
|
||||||
)
|
|
||||||
continue
|
continue
|
||||||
|
markdown_image, local_image_path = build_image_markdown(str(item), prefix)
|
||||||
img_path = str(img_path)
|
markdown_images.append(markdown_image)
|
||||||
if img_path.strip().startswith("!["):
|
if local_image_path:
|
||||||
local_image_path = resolve_local_image_path(img_path)
|
local_image_paths.append(local_image_path)
|
||||||
ins["img_path"] = build_image_markdown_with_ocr(img_path, local_image_path)
|
ins["img_path"] = "\n".join(markdown_images)
|
||||||
continue
|
|
||||||
|
|
||||||
p = PurePosixPath(img_path)
|
|
||||||
# 如果以 images/ 开头,去掉第一级目录
|
|
||||||
if p.parts and p.parts[0] == "images":
|
|
||||||
relative_path = str(PurePosixPath(*p.parts[1:]))
|
|
||||||
markdown_image = f""
|
|
||||||
local_image_path = resolve_local_image_path(relative_path)
|
|
||||||
ins["img_path"] = build_image_markdown_with_ocr(markdown_image, local_image_path)
|
|
||||||
else:
|
else:
|
||||||
# 否则保留原路径(或按需处理)
|
markdown_image, local_image_path = build_image_markdown(str(img_path), prefix)
|
||||||
markdown_image = f""
|
ins["img_path"] = markdown_image
|
||||||
local_image_path = resolve_local_image_path(img_path)
|
if local_image_path:
|
||||||
ins["img_path"] = build_image_markdown_with_ocr(markdown_image, local_image_path)
|
local_image_paths.append(local_image_path)
|
||||||
|
|
||||||
|
ocr_texts = await asyncio.gather(
|
||||||
|
*(get_image_ocr_text(path) for path in local_image_paths)
|
||||||
|
)
|
||||||
|
append_caption_ocr(ins, ocr_texts)
|
||||||
return data
|
return data
|
||||||
|
|
||||||
def filter_content_bbox(data, max_length=3000):
|
def filter_content_bbox(data, max_length=3000):
|
||||||
@ -1034,45 +1132,15 @@ def filter_content_bbox(data, max_length=3000):
|
|||||||
if "type" not in ins:
|
if "type" not in ins:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
added_content = ""
|
|
||||||
bbox = ins.get("bbox", [])
|
bbox = ins.get("bbox", [])
|
||||||
page_idx = ins.get("page_idx", -1)
|
page_idx = ins.get("page_idx", -1)
|
||||||
|
|
||||||
if ins["type"] == "text":
|
if not is_supported_chunk_record(ins):
|
||||||
text_content = ins.get("text", "")
|
continue
|
||||||
added_content = text_content + "\n"
|
|
||||||
|
|
||||||
elif ins["type"] == "table":
|
added_content = record_to_chunk_text(ins)
|
||||||
table_caption = ins.get("table_caption")
|
if added_content.strip():
|
||||||
table_body = ins.get("table_body")
|
added_content += "\n"
|
||||||
|
|
||||||
caption_str = ""
|
|
||||||
if table_caption:
|
|
||||||
if isinstance(table_caption, list):
|
|
||||||
caption_str = "\n".join(str(c) for c in table_caption)
|
|
||||||
else:
|
|
||||||
caption_str = str(table_caption)
|
|
||||||
|
|
||||||
added_content = caption_str + "\n" + (table_body or "") + "\n"
|
|
||||||
|
|
||||||
elif ins["type"] == "image":
|
|
||||||
image_caption = ins.get("image_caption", "")
|
|
||||||
img_path = ins.get("img_path", "")
|
|
||||||
|
|
||||||
if isinstance(image_caption, list):
|
|
||||||
image_caption = "\n".join(str(item) for item in image_caption)
|
|
||||||
elif not isinstance(image_caption, str):
|
|
||||||
image_caption = str(image_caption)
|
|
||||||
|
|
||||||
if isinstance(img_path, list):
|
|
||||||
img_path = "\n".join(str(p) for p in img_path)
|
|
||||||
elif not isinstance(img_path, str):
|
|
||||||
img_path = str(img_path)
|
|
||||||
|
|
||||||
if image_caption or img_path:
|
|
||||||
added_content = image_caption + "\n" + img_path + "\n"
|
|
||||||
else:
|
|
||||||
continue # 跳过空图像项
|
|
||||||
else:
|
else:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@ -1319,7 +1387,8 @@ async def process_pdf_file(pdf_path: Path, image_prefix: str,filename:str) -> Op
|
|||||||
if doc_result is not None:
|
if doc_result is not None:
|
||||||
content_list, images = doc_result
|
content_list, images = doc_result
|
||||||
if content_list and images:
|
if content_list and images:
|
||||||
replaced_content_list = data_replace(content_list, image_prefix)
|
replaced_content_textlevel = reset_textlevel(content_list)
|
||||||
|
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
|
||||||
slices = get_chunk_bbox(replaced_content_list)
|
slices = get_chunk_bbox(replaced_content_list)
|
||||||
slices_check = chunk_check(slices, 8000)
|
slices_check = chunk_check(slices, 8000)
|
||||||
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
|
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
|
||||||
@ -1342,7 +1411,8 @@ async def process_pdf_file(pdf_path: Path, image_prefix: str,filename:str) -> Op
|
|||||||
data = result.get("data", {})
|
data = result.get("data", {})
|
||||||
content_list = data.get("content_list", [])
|
content_list = data.get("content_list", [])
|
||||||
images = data.get("images", {})
|
images = data.get("images", {})
|
||||||
replaced_content_list = data_replace(content_list, image_prefix)
|
replaced_content_textlevel = reset_textlevel(content_list)
|
||||||
|
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
|
||||||
slices = get_chunk_bbox(replaced_content_list)
|
slices = get_chunk_bbox(replaced_content_list)
|
||||||
slices_check = chunk_check(slices, 8000)
|
slices_check = chunk_check(slices, 8000)
|
||||||
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
|
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
|
||||||
@ -1367,7 +1437,8 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) -
|
|||||||
if doc_result is not None:
|
if doc_result is not None:
|
||||||
content_list, images = doc_result
|
content_list, images = doc_result
|
||||||
if content_list and images:
|
if content_list and images:
|
||||||
replaced_content_list = data_replace(content_list, image_prefix)
|
replaced_content_textlevel = reset_textlevel(content_list)
|
||||||
|
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
|
||||||
slices = get_chunk_bbox(replaced_content_list)
|
slices = get_chunk_bbox(replaced_content_list)
|
||||||
slices_check = chunk_check(slices, 8000)
|
slices_check = chunk_check(slices, 8000)
|
||||||
return {"slices": slices_check, "images": images}
|
return {"slices": slices_check, "images": images}
|
||||||
@ -1397,7 +1468,8 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) -
|
|||||||
images = data.get("images", {})
|
images = data.get("images", {})
|
||||||
|
|
||||||
# 后续的数据替换与切片处理逻辑
|
# 后续的数据替换与切片处理逻辑
|
||||||
replaced_content_list = data_replace(content_list, image_prefix)
|
replaced_content_textlevel = reset_textlevel(content_list)
|
||||||
|
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
|
||||||
slices = get_chunk_bbox(replaced_content_list)
|
slices = get_chunk_bbox(replaced_content_list)
|
||||||
slices_check = chunk_check(slices, 8000)
|
slices_check = chunk_check(slices, 8000)
|
||||||
|
|
||||||
@ -1415,7 +1487,7 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) -
|
|||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|
||||||
filepath = r"E:\ZKYNLP\Hjunproject\project0506\kgrag\122-06A0014-B01001_发动机-维修手册_content_list.json"
|
filepath = r"E:\ZKYNLP\Hjunproject\project0506\kgrag\舰船抗沉损管训练仿真系统研究_content_list.json"
|
||||||
try:
|
try:
|
||||||
with open(filepath, 'r', encoding='utf-8') as file:
|
with open(filepath, 'r', encoding='utf-8') as file:
|
||||||
data = json.load(file)
|
data = json.load(file)
|
||||||
@ -1425,12 +1497,12 @@ if __name__ == "__main__":
|
|||||||
# print(len(data))
|
# print(len(data))
|
||||||
start_chars = ["1", "2", "3", "4", "5", "6", "7", "8", "9", "0", "第", "十",
|
start_chars = ["1", "2", "3", "4", "5", "6", "7", "8", "9", "0", "第", "十",
|
||||||
"一", "二", "三", "四", "五", "六", "七", "八", "九"]
|
"一", "二", "三", "四", "五", "六", "七", "八", "九"]
|
||||||
|
replaced_content_textlevel = reset_textlevel(data)
|
||||||
data = data_replace(data=data,prefix="/api/v1/knowledge/files/images/")
|
data = asyncio.run(data_replace(data=replaced_content_textlevel,prefix="/api/v1/knowledge/files/images/"))
|
||||||
slices = get_chunk_bbox(data=data)
|
slices = get_chunk_bbox(data=data)
|
||||||
slices_1 =chunk_check(slices, 8000)
|
slices_1 =chunk_check(slices, 8000)
|
||||||
slices_2 = merge_short_slices(slices_1, min_length=30) # ← 新增这一行
|
slices_2 = merge_short_slices(slices_1, min_length=30) # ← 新增这一行
|
||||||
|
|
||||||
|
|
||||||
for ins in slices_2:
|
for ins in slices_2[:30]:
|
||||||
print(ins)
|
print(ins)
|
||||||
|
|||||||
184
resetlevel.py
Normal file
184
resetlevel.py
Normal file
@ -0,0 +1,184 @@
|
|||||||
|
import copy
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
|
CN_NUM = "\u4e00\u4e8c\u4e09\u56db\u4e94\u516d\u4e03\u516b\u4e5d\u5341\u767e\u5343\u4e07\u3007\u96f6\u4e24"
|
||||||
|
|
||||||
|
CHAPTER_PATTERN = re.compile(rf"^\s*\u7b2c\s*([{CN_NUM}\d]+)\s*\u7ae0")
|
||||||
|
CN_LEVEL1_PATTERN = re.compile(rf"^\s*([{CN_NUM}]+)(?=\s|[\u3001\u3002\uff0e.])")
|
||||||
|
|
||||||
|
CASE_LIKE_PATTERN = re.compile(r"^\s*\d+(?:\.|\uff0e)(?!\d|\s|$)")
|
||||||
|
ARABIC_PATTERN = re.compile(
|
||||||
|
r"^\s*(\d+(?:\.\d+)*)(?=\s|[\u4e00-\u9fff\u3001\u3002\uff0e]|$|\.(?:\s|$))"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def has_valid_text_level(item: dict) -> bool:
|
||||||
|
return "text_level" in item and item.get("text_level") not in (None, "")
|
||||||
|
|
||||||
|
|
||||||
|
def parse_heading(text: str):
|
||||||
|
text = text.strip()
|
||||||
|
if not text:
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
if CASE_LIKE_PATTERN.match(text):
|
||||||
|
return "case_like", None
|
||||||
|
|
||||||
|
if CHAPTER_PATTERN.match(text):
|
||||||
|
return "level1", None
|
||||||
|
|
||||||
|
if CN_LEVEL1_PATTERN.match(text):
|
||||||
|
return "level1", None
|
||||||
|
|
||||||
|
match = ARABIC_PATTERN.match(text)
|
||||||
|
if match:
|
||||||
|
return "arabic", match.group(1)
|
||||||
|
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
|
||||||
|
def arabic_depth(number: str) -> int:
|
||||||
|
return number.count(".") + 1
|
||||||
|
|
||||||
|
|
||||||
|
def parent_number(number: str) -> str | None:
|
||||||
|
parts = number.split(".")
|
||||||
|
if len(parts) <= 1:
|
||||||
|
return None
|
||||||
|
return ".".join(parts[:-1])
|
||||||
|
|
||||||
|
|
||||||
|
def collect_headings(data: list) -> list:
|
||||||
|
headings = []
|
||||||
|
|
||||||
|
for item in data:
|
||||||
|
if not isinstance(item, dict):
|
||||||
|
continue
|
||||||
|
if item.get("type") != "text":
|
||||||
|
continue
|
||||||
|
if not has_valid_text_level(item):
|
||||||
|
continue
|
||||||
|
|
||||||
|
kind, number = parse_heading(item.get("text", ""))
|
||||||
|
if kind in {"level1", "arabic", "case_like"}:
|
||||||
|
headings.append({"kind": kind, "number": number})
|
||||||
|
|
||||||
|
return headings
|
||||||
|
|
||||||
|
|
||||||
|
def is_valid_arabic_hierarchy(numbers: list[str], allow_missing_depth1_parent: bool) -> bool:
|
||||||
|
seen = set()
|
||||||
|
has_depth2 = False
|
||||||
|
has_depth3 = False
|
||||||
|
|
||||||
|
for number in numbers:
|
||||||
|
depth = arabic_depth(number)
|
||||||
|
if depth >= 2:
|
||||||
|
has_depth2 = True
|
||||||
|
if depth >= 3:
|
||||||
|
has_depth3 = True
|
||||||
|
|
||||||
|
parent = parent_number(number)
|
||||||
|
if parent is not None:
|
||||||
|
parent_depth = arabic_depth(parent)
|
||||||
|
parent_is_missing_depth1 = allow_missing_depth1_parent and parent_depth == 1
|
||||||
|
if parent not in seen and not parent_is_missing_depth1:
|
||||||
|
return False
|
||||||
|
|
||||||
|
seen.add(number)
|
||||||
|
|
||||||
|
return has_depth2 and has_depth3
|
||||||
|
|
||||||
|
|
||||||
|
def detect_valid_mode(headings: list):
|
||||||
|
if not headings:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if any(heading["kind"] == "case_like" for heading in headings):
|
||||||
|
return None
|
||||||
|
|
||||||
|
first = headings[0]
|
||||||
|
|
||||||
|
if first["kind"] == "arabic":
|
||||||
|
numbers = [heading["number"] for heading in headings if heading["kind"] == "arabic"]
|
||||||
|
if numbers and arabic_depth(numbers[0]) == 1:
|
||||||
|
if is_valid_arabic_hierarchy(numbers, allow_missing_depth1_parent=False):
|
||||||
|
return "A"
|
||||||
|
return None
|
||||||
|
|
||||||
|
if first["kind"] == "level1":
|
||||||
|
numbers = [heading["number"] for heading in headings[1:] if heading["kind"] == "arabic"]
|
||||||
|
if not numbers:
|
||||||
|
return None
|
||||||
|
|
||||||
|
first_depth = arabic_depth(numbers[0])
|
||||||
|
if first_depth == 2 and is_valid_arabic_hierarchy(numbers, allow_missing_depth1_parent=True):
|
||||||
|
return "B"
|
||||||
|
|
||||||
|
if first_depth == 1 and is_valid_arabic_hierarchy(numbers, allow_missing_depth1_parent=False):
|
||||||
|
return "C"
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
JsonContent = Any
|
||||||
|
|
||||||
|
|
||||||
|
def reset_textlevel(json_content: JsonContent) -> JsonContent:
|
||||||
|
"""
|
||||||
|
输入 JSON 内容,返回更新 text_level 后的新 JSON 内容。
|
||||||
|
|
||||||
|
json_content 可以是已经 json.load/json.loads 后的 Python 对象,
|
||||||
|
也可以是 JSON 字符串。函数不会修改原始入参。
|
||||||
|
"""
|
||||||
|
if isinstance(json_content, str):
|
||||||
|
data = json.loads(json_content)
|
||||||
|
else:
|
||||||
|
data = copy.deepcopy(json_content)
|
||||||
|
|
||||||
|
if not isinstance(data, list):
|
||||||
|
return data
|
||||||
|
|
||||||
|
headings = collect_headings(data)
|
||||||
|
mode = detect_valid_mode(headings)
|
||||||
|
|
||||||
|
if mode is None:
|
||||||
|
return data
|
||||||
|
|
||||||
|
for item in data:
|
||||||
|
if not isinstance(item, dict):
|
||||||
|
continue
|
||||||
|
if item.get("type") != "text":
|
||||||
|
continue
|
||||||
|
if not has_valid_text_level(item):
|
||||||
|
continue
|
||||||
|
|
||||||
|
kind, number = parse_heading(item.get("text", ""))
|
||||||
|
if kind == "level1":
|
||||||
|
item["text_level"] = 1
|
||||||
|
continue
|
||||||
|
if kind != "arabic":
|
||||||
|
continue
|
||||||
|
|
||||||
|
depth = arabic_depth(number)
|
||||||
|
item["text_level"] = depth if mode in {"A", "B"} else depth + 1
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
json_file_path = r'E:\ZKYNLP\Hjunproject\project0506\kgrag\163-06A0014-B01001_发动机-维修手册_content_list.json'
|
||||||
|
with open(json_file_path, "r", encoding="utf-8") as f:
|
||||||
|
json_content = json.load(f)
|
||||||
|
result= reset_textlevel(json_content)
|
||||||
|
|
||||||
|
for ins in result[:30]:
|
||||||
|
print(ins)
|
||||||
|
# print(json.dumps(result, ensure_ascii=False, indent=2))
|
||||||
Loading…
x
Reference in New Issue
Block a user