kgrag/app_pkg/examples/markdown_splits.py
2026-07-29 18:10:19 +08:00

177 lines
6.0 KiB
Python

from langchain.text_splitter import MarkdownTextSplitter,MarkdownHeaderTextSplitter,RecursiveCharacterTextSplitter
from md_header_convert import convert_numbered_headings
import re
import json
text_splitter = RecursiveCharacterTextSplitter(
chunk_size=256,
chunk_overlap=50,
length_function=len,
separators=[" "]
)
headers_to_split_on = [("#", "Header_1"), ("##", "Header_2"), ("###", "Header_3")]
markdown_splitter = MarkdownHeaderTextSplitter(headers_to_split_on=headers_to_split_on)
def parse_markdown(md_text):
md_text = convert_numbered_headings(md_text)
lines = md_text.strip().split("\n")
results = []
current_h1, current_h2, current_h3 = None, None, None
current_text = []
def add_entry():
if current_h1 or current_h2 or current_h3:
results.append({
"Header_1": current_h1,
"Header_2": current_h2,
"Header_3": current_h3,
"text": "\n".join(current_text).strip()
})
current_text.clear()
for line in lines:
line = line.strip()
header_match = re.match(r'^(#{1,3})\s*(.*)', line)
if header_match:
add_entry()
level = len(header_match.group(1))
title = header_match.group(2).strip()
if level == 1:
current_h1, current_h2, current_h3 = title, None, None
elif level == 2:
current_h2, current_h3 = title, None
elif level == 3:
current_h3 = title
else:
current_text.append(line)
add_entry()
return results
# return json.dumps(results, indent=4, ensure_ascii=False)
def header_postprocess(md_header_json):
for md in md_header_json:
# print(md)
if not md["Header_2"] and not md["Header_3"]:
if md["text"]:
continue
else:
md["text"] = md["Header_1"]
if md["Header_2"] and not md["Header_3"]:
if md["text"]:
md["text"] = md["Header_2"] + md["text"]
md["Header_2"] = None
else:
md["text"] = md["Header_2"]
md["Header_2"] = None
if md["Header_2"] and md["Header_3"]:
if not md["text"]:
md["text"] = md["Header_3"]
md["Header_3"] = None
else:
md["text"] = md["Header_3"] + md["text"]
md["Header_3"] = None
return md_header_json
# text_splitter = MarkdownTextSplitter(chunk_size=512,chunk_overlap=128)
# print(text)
# print("####################")
def langchain_split(text):
text = convert_numbered_headings(text)
# print(text)
md_header_splits = markdown_splitter.split_text(text)
results = []
for i, doc in enumerate(md_header_splits):
print("-------------------------------------------------------")
print(f"Document {i+1}:")
print("text")
print(doc.page_content)
if doc.metadata:
header_1 = doc.metadata["Header_1"] if "Header_1" in doc.metadata.keys() else None
header_2 = doc.metadata["Header_2"] if "Header_2" in doc.metadata.keys() else None
header_3 = doc.metadata["Header_3"] if "Header_3" in doc.metadata.keys() else None
text_ = doc.page_content
# print("Header_1:",header_1)
# print("Header_2:",header_2)
# print("Header_3:",header_3)
results.append({"Header_1":header_1,"Header_2":header_2,"Header_3":header_3,"text":text_})
# print("\n")
return results
# md_header_json = header_postprocess(results)
# for md in md_header_json:
# print(md)
# print("======================================")
def split_regular(text):
#自定义分解
md_header_json = parse_markdown(text)
# print(md_header_json)
#合并处理
md_header_json = header_postprocess(md_header_json)
# print(md_header_json)
results = []
for i ,md in enumerate(md_header_json):
if md["text"]:
header_1 = md["Header_1"]
header_2 = md["Header_2"]
header_3 = md["Header_3"]
text = md["text"]
# print("Header_1:",header_1)
# print("Header_2:",header_2)
# print("Header_3:",header_3)
# print("text:",text_)
# print("======================================")
texts = text_splitter.split_text(text)
for j,txt in enumerate(texts):
# print(j)
# print(texts)
# print("======================================")
results.append({"id":j+i+1,"Header_1":header_1,"Header_2":header_2,"Header_3":header_3,"text":txt})
return results
# results = []
# for _, i in enumerate(md_header_splits):
# if i.metadata:
# header_1 = i.metadata["Header_1"] if "Header_1" in i.metadata.keys() else None
# header_2 = i.metadata["Header_2"] if "Header_2" in i.metadata.keys() else None
# header_3 = i.metadata["Header_3"] if "Header_3" in i.metadata.keys() else None
# text_ = i.page_content
# print("Header_1:",header_1)
# print("Header_2:",header_2)
# print("Header_3:",header_3)
# # texts = text_splitter.split_text(text)
# # for j in texts:
# # print(j)
# # print(texts)
# print("======================================")
# # results.append({"id":i,"content":texts})
# for res in results:
# print(res["id"])
# print(res["content"])
# print("======================================")
if __name__ == "__main__":
with open("邱茜茜.md","r",encoding="utf-8") as f:
text = f.read()
results = split_regular(text)
for res in results:
print("======================================")
for key, value in res.items():
print(f"{key}: {value}")
print("\n")