177 lines
6.0 KiB
Python
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")
|
|
|