更新 ocr 模型

This commit is contained in:
Wenjie Zhang 2025-03-08 20:26:00 +08:00
parent e3ee532aca
commit 8132ac2e9b
13 changed files with 235 additions and 173 deletions

2
.gitignore vendored
View File

@ -38,3 +38,5 @@ notebooks
graphrag
docker/volumes
.cursorrules
/models

View File

@ -8,7 +8,7 @@ services:
volumes:
- ../src:/app/src
- ../saves_dev:/app/saves
# - ${MODEL_DIR}:${MODEL_DIR} # 如果出现 undefined volume MODEL_DIR: invalid compose project请添加此环境变量或者注释此行
# - ${MODEL_DIR}:/models # 如果出现 undefined volume MODEL_DIR: invalid compose project请添加此环境变量或者注释此行
ports:
- "5050:5050"
depends_on:
@ -23,7 +23,7 @@ services:
- NEO4J_USERNAME=neo4j
- NEO4J_PASSWORD=0123456789
- MILVUS_URI=http://milvus:19530
- MODEL_DIR=${MODEL_DIR} # 优先级高于 .env 中的 MODEL_DIR
- MODEL_DIR=/models # 优先级高于 .env 中的 MODEL_DIR
- RUNNING_IN_DOCKER=true
command: uvicorn src.main:app --host 0.0.0.0 --port 5050 --reload

View File

@ -25,4 +25,5 @@ docx2txt
uvicorn[standard]
fastapi
python-multipart
tavily-python
tavily-python
rapidocr_onnxruntime

View File

@ -0,0 +1,5 @@
#! /bin/sh
source src/.env
# OCR 模型
huggingface-cli download SWHL/RapidOCR --local-dir ${MODEL_DIR}/SWHL/RapidOCR

3
src/__init__.py Normal file
View File

@ -0,0 +1,3 @@
from dotenv import load_dotenv
load_dotenv("src/.env")

View File

@ -1,7 +1,6 @@
import os
import json
import time
from src.plugins import pdf2txt
from src.utils import hashstr, logger, is_text_pdf
from src.models.embedding import get_embedding_model
@ -177,12 +176,8 @@ class DataBaseManager:
raise NotImplementedError("Directory not supported now!")
if file.endswith(".pdf"):
if is_text_pdf(file):
from src.core.filereader import pdfreader
return pdfreader(file)
else:
from src.plugins import pdf2txt
return pdf2txt(file, return_text=True)
from src.plugins import ocr
return ocr.process_pdf(file)
elif file.endswith(".txt") or file.endswith(".md"):
from src.core.filereader import plainreader

View File

@ -124,16 +124,6 @@ class GraphDatabase:
with self.driver.session() as session:
session.execute_write(create, triples)
# def pdf_file_add_entity(self, file_path, output_path, kgdb_name='neo4j'):
# self.use_database(kgdb_name) # 切换到指定数据库
# text_path = pdf2txt(file_path)
# global UIE_MODEL
# if UIE_MODEL is None:
# UIE_MODEL = OneKE()
# triples_path = UIE_MODEL.processing_text_to_kg(text_path, output_path)
# self.jsonl_file_add_entity(triples_path)
# return kgdb_name
def txt_add_vector_entity(self, triples, kgdb_name='neo4j'):
"""添加实体三元组"""
self.use_database(kgdb_name)

View File

@ -1,7 +1,4 @@
import uvicorn
from dotenv import load_dotenv
load_dotenv("src/.env")
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware

View File

@ -1,2 +1,5 @@
from src.plugins.oneke import *
from src.plugins.pdf2txt import *
from src.plugins._ocr import OCRPlugin
ocr = OCRPlugin()
__all__ = ["ocr"]

210
src/plugins/_ocr.py Normal file
View File

@ -0,0 +1,210 @@
import os
import uuid
from pathlib import Path
from argparse import ArgumentParser
import fitz # fitz就是pip install PyMuPDF
import numpy as np # Added import for numpy
from PIL import Image
from tqdm import tqdm
from rapidocr_onnxruntime import RapidOCR
from src.utils import logger, is_text_pdf
GOLBAL_STATE = {}
class OCRPlugin:
"""OCR 插件"""
def __init__(self, **kwargs):
self.ocr = None
self.det_box_thresh = kwargs.get('det_box_thresh', 0.3)
def load_model(self):
"""加载 OCR 模型"""
logger.info(f"加载 OCR 模型,仅在第一次调用时加载")
model_dir = os.path.join(os.getenv("MODEL_DIR", ""), "SWHL/RapidOCR")
det_model_dir = os.path.join(model_dir, "PP-OCRv4/ch_PP-OCRv4_det_infer.onnx")
rec_model_dir = os.path.join(model_dir, "PP-OCRv4/ch_PP-OCRv4_rec_infer.onnx")
assert os.path.exists(model_dir), (
f"模型文件不存在,请下载 SWHL/RapidOCR 到 {model_dir}"
"并确认是否在 docker-compose.dev.yml 中添加 MODEL_DIR 环境变量"
)
self.ocr = RapidOCR(det_box_thresh=0.3, det_model_path=det_model_dir, rec_model_path=rec_model_dir)
logger.info(f"OCR Plugin for det_box_thresh = {self.det_box_thresh} loaded.")
def process_image(self, image):
"""
对单张图像执行OCR并提取文本
Args:
image: 图像数据支持多种格式
- str: 图像文件路径
- PIL.Image: PIL图像对象
- numpy.ndarray: numpy图像数组
Returns:
str: 提取的文本内容
"""
# 确保模型已加载
if self.ocr is None:
self.load_model()
# 处理不同类型的输入图像
try:
if isinstance(image, str):
# 图像路径直接传递给OCR处理
image_path = image
is_temp_file = False
else:
# 创建临时文件
is_temp_file = True
image_path = self._create_temp_image_file(image)
# 执行 OCR
result, _ = self.ocr(image_path)
# 清理临时文件
if is_temp_file and os.path.exists(image_path):
os.remove(image_path)
# 提取文本
if result:
text = '\n'.join([line[1] for line in result])
return text
else:
logger.warning(f"OCR未能识别出文本内容")
return ""
except Exception as e:
logger.error(f"OCR处理失败: {str(e)}")
raise
def _create_temp_image_file(self, image):
"""
将图像数据保存为临时文件
Args:
image: PIL.Image或numpy.ndarray格式的图像数据
Returns:
str: 临时文件路径
"""
# 为临时文件创建目录(如果不存在)
tmp_dir = os.path.join(os.getcwd(), 'tmp')
os.makedirs(tmp_dir, exist_ok=True)
# 生成临时文件路径
temp_filename = f'ocr_temp_{uuid.uuid4().hex[:8]}.png'
image_path = os.path.join(tmp_dir, temp_filename)
# 根据图像类型保存文件
if isinstance(image, Image.Image):
# 保存PIL图像对象到临时文件
image.save(image_path)
elif isinstance(image, np.ndarray):
# 将numpy数组转换为PIL图像并保存
Image.fromarray(image).save(image_path)
else:
raise ValueError("不支持的图像类型必须是PIL.Image或numpy数组")
return image_path
def process_pdf(self, pdf_path):
"""
处理PDF文件并提取文本
:param pdf_path: PDF文件路径
:return: 提取的文本
"""
if self.ocr is None:
self.load_model()
if not os.path.exists(pdf_path):
raise FileNotFoundError(f"PDF file not found: {pdf_path}")
try:
# 检查是否为文本PDF
if is_text_pdf(pdf_path):
logger.info(f"PDF file is text, use llama_index.readers.file to read")
return pdfreader(pdf_path)
# 将PDF转换为图像
filename = os.path.basename(pdf_path).split('.')[0]
output_dir = os.path.join('saves', 'data', 'pdf2txt', filename)
os.makedirs(output_dir, exist_ok=True)
images = self.convert_imgs(pdf_path, output_dir)
# 处理每个图像并合并文本
all_text = []
for img_path in tqdm(images, desc='to txt', ncols=100):
text = self.process_image(img_path)
all_text.append(text)
return '\n\n'.join(all_text)
except Exception as e:
logger.error(f"PDF processing error: {str(e)}")
return ""
def convert_imgs(self, pdf_path, output_dir):
imgs = []
img_dir = os.path.join(output_dir, 'imgs')
if not os.path.exists(img_dir):
os.makedirs(img_dir)
pdfDoc = fitz.open(pdf_path)
totalPage = pdfDoc.page_count
for pg in tqdm(range(totalPage), desc='to imgs', ncols=100):
page = pdfDoc[pg]
rotate = int(0)
zoom_x = 2
zoom_y = 2
mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate)
pix = page.get_pixmap(matrix=mat, alpha=False)
img_filename = os.path.join(img_dir, f'images_{pg+1}.png')
pix.save(img_filename) # os.sep
imgs.append(img_filename)
else:
img_names = sorted(os.listdir(img_dir))
imgs = [os.path.join(img_dir, img_name) for img_name in img_names]
return imgs
def get_state(task_id):
return GOLBAL_STATE.get(task_id, {})
def pdfreader(file_path):
"""读取PDF文件并返回text文本"""
assert os.path.exists(file_path), "File not found"
assert file_path.endswith(".pdf"), "File format not supported"
from llama_index.readers.file import PDFReader
doc = PDFReader().load_data(file=Path(file_path))
# 简单的拼接起来之后返回纯文本
text = "\n\n".join([d.get_content() for d in doc])
return text
def plainreader(file_path):
"""读取普通文本文件并返回text文本"""
assert os.path.exists(file_path), "File not found"
with open(file_path, "r") as f:
text = f.read()
return text
if __name__ == "__main__":
parser = ArgumentParser()
parser.add_argument('--pdf-path', type=str, required=True, help='Path to the PDF file')
parser.add_argument('--return-text', action='store_true', help='Return the extracted text')
args = parser.parse_args()
ocr = OCRPlugin()
text = ocr.process_pdf(args.pdf_path)
print(text)

View File

@ -1,144 +0,0 @@
import os
import uuid
import fitz # fitz就是pip install PyMuPDF
from PIL import Image
from copy import deepcopy
from tqdm import tqdm
from pathlib import Path
from argparse import ArgumentParser
from src.utils import logger, is_text_pdf
import numpy as np # Added import for numpy
GOLBAL_STATE = {}
def pdf2txt(pdf_path, return_text=False):
if not os.path.exists(pdf_path):
raise FileNotFoundError(f"File not found: {pdf_path}")
if is_text_pdf(pdf_path):
return pdfreader(pdf_path)
# Importing these modules here to avoid unnecessary imports in other files
from paddleocr import PPStructure, save_structure_res
from paddleocr.ppstructure.recovery.recovery_to_doc import sorted_layout_boxes, convert_info_docx
filename = os.path.basename(pdf_path).split('.')[0]
output_dir = os.path.join('saves', 'data', 'pdf2txt', filename)
os.makedirs(output_dir, exist_ok=True)
table_engine = PPStructure(recovery=True, lang='ch')
imgs = convert_imgs(pdf_path, output_dir)
respath = os.path.join(output_dir, f'{filename}.txt')
global GOLBAL_STATE
task_id = str(uuid.uuid4())
if task_id in GOLBAL_STATE:
logger.info(f"Reusing previous state for process {task_id}")
return GOLBAL_STATE[task_id]
else:
logger.info(f"Creating new state for process {task_id}")
GOLBAL_STATE[task_id] = {
'pdf_path': pdf_path,
'return_text': return_text,
'output_dir': output_dir,
'status': 'in-progress',
'total': len(imgs),
'progress': 0
}
text = []
for img_name in tqdm(imgs, desc='to txt', ncols=100):
img = Image.open(img_name)
img = np.array(img) # 如果需要将图像转换为numpy数组
result = table_engine(img)
GOLBAL_STATE[task_id]['progress'] += 1
save_structure_res(result, output_dir, "structure_result")
res = sorted_layout_boxes(result, img.shape[1])
try:
convert_info_docx(img, res, output_dir, "info")
except Exception as e:
print(f"Exception occurred while converting to DOCX: {e}")
continue # Skip this image if conversion fails
for line in res:
line.pop('img') # okk
for pra in line['res']:
if isinstance(pra, dict) and 'text' in pra:
text.append(pra['text'])
else:
continue # 如果不是字典或者缺少 'text' 键,跳过当前循环
text.append('\n')
whole_text = ''.join(text)
with open(respath, 'w', encoding='utf-8') as f:
f.write(whole_text)
logger.info(f"Extracted text saved to {respath}")
GOLBAL_STATE[task_id]['status'] = 'completed'
if return_text:
return whole_text
return respath
def convert_imgs(pdf_path, output_dir):
imgs = []
img_dir = os.path.join(output_dir, 'imgs')
if not os.path.exists(img_dir):
os.makedirs(img_dir)
pdfDoc = fitz.open(pdf_path)
totalPage = pdfDoc.page_count
for pg in tqdm(range(totalPage), desc='to imgs', ncols=100):
page = pdfDoc[pg]
rotate = int(0)
zoom_x = 2
zoom_y = 2
mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate)
pix = page.get_pixmap(matrix=mat, alpha=False)
img_filename = os.path.join(img_dir, f'images_{pg+1}.png')
pix.save(img_filename) # os.sep
imgs.append(img_filename)
else:
img_names = sorted(os.listdir(img_dir))
imgs = [os.path.join(img_dir, img_name) for img_name in img_names]
return imgs
def get_state(task_id):
return GOLBAL_STATE.get(task_id, {})
def pdfreader(file_path):
"""读取PDF文件并返回text文本"""
assert os.path.exists(file_path), "File not found"
assert file_path.endswith(".pdf"), "File format not supported"
from llama_index.readers.file import PDFReader
doc = PDFReader().load_data(file=Path(file_path))
# 简单的拼接起来之后返回纯文本
text = "\n\n".join([d.get_content() for d in doc])
return text
def plainreader(file_path):
"""读取普通文本文件并返回text文本"""
assert os.path.exists(file_path), "File not found"
with open(file_path, "r") as f:
text = f.read()
return text
if __name__ == "__main__":
parser = ArgumentParser()
parser.add_argument('--pdf-path', type=str, required=True, help='Path to the PDF file')
parser.add_argument('--return-text', action='store_true', help='Return the extracted text')
args = parser.parse_args()
pdf2txt(args.pdf_path, args.return_text)

View File

@ -45,6 +45,6 @@ async def text_chunking(text: str = Body(...), params: Dict[str, Any] = Body(...
@tool.post("/pdf2txt")
async def handle_pdf2txt(file: str = Body(...)):
from src.plugins import pdf2txt
text = pdf2txt(file, return_text=True)
from src.plugins import ocr
text = ocr.process_pdf(file)
return {"text": text}

View File

@ -189,7 +189,7 @@
</div>
<div class="setting" v-if="state.windowWidth <= 520 || state.section ==='path'">
<h3>本地模型配置</h3>
<p>如果是 Docker 启动务必确保在环境变量中设置了 MODEL_DIR建议是绝对路径并设置volumes 映射</p>
<p>如果是 Docker 启动务必确保在 docker-compose.dev.yaml 中添加volumes 映射</p>
<TableConfigComponent
:config="configStore.config?.model_local_paths"
@update:config="handleModelLocalPathsUpdate"
@ -508,7 +508,7 @@ const sendRestart = () => {
width: 100%;
height: 100%;
}
&.available {
filter: grayscale(0%);
}