ocr_api.py

ocr_api.py — PaddleOCR 封装类

初始化模型、长图智能切分(OpenCV 连通组件检测判断切线是否切到文字,切到了就上移)、结果缓存、GIF 兼容处理。切图算法详解见 api 主篇

代码

# -*- coding: utf-8 -*-
# @Time    : 2023/10/23 15:40
# @Author  : chang
'''
飞将OCR识别(基于PaddleOCR)
官方文档:https://github.com/PaddlePaddle/PaddleOCR/blob/release/2.7/doc/doc_ch/quickstart.md
功能:对长图自动切分、GIF自动转PNG、OCR文字识别、结果缓存、自动按行排版输出纯文本
'''

# 导入PaddleOCR核心识别类
from paddleocr import PaddleOCR
# 导入PIL图片处理库,用于图片打开、裁剪、格式转换
from PIL import Image
# 导入系统路径、文件操作库
import os
# 导入时间库,用于生成时间戳
import time
# 导入OpenCV,用于图片二值化、连通组件分析(判断是否有文字)
import cv2
# 导入数值计算库,用于数组处理
import numpy as np

# 导入系统库,添加模块搜索路径
import sys
sys.path.append('../')  # 添加上级目录到Python路径,用于导入自定义工具

# 导入自定义工具:GIF转PNG函数
from utils_img import convert_gif_png
# 导入自定义日志工具
from utils import ner_logger

# PaddleOCR目前支持的多语言语种可以通过修改lang参数进行切换
# 例如`ch`(中文), `en`(英文), `fr`(法语), `german`, `korean`, `japan`

class Paddle_OCR():
    """
    PaddleOCR封装类
    功能:初始化模型、长图智能切分、单图/长图OCR识别、结果缓存、GIF兼容处理
    """
    def __init__(self):
        """
        构造函数:初始化OCR模型
        use_angle_cls=True:开启方向分类器,自动修正倾斜图片
        lang="ch":使用中文模型
        模型只会加载一次到内存,后续重复使用
        """
        self.OCR = PaddleOCR(use_angle_cls=True, lang="ch")  
        # 定义时间戳生成函数:毫秒级时间戳
        self.TIMENUMS=lambda :int(time.time()*1000)

    def img_cut(self,img_path,img_savepath=None,part_height=1000):#640
        '''
        智能切图函数:解决超长图片OCR识别效果差的问题
        智能判断文字位置,在无文字区域切分,保证文字完整性
        :param img_path: 原始图片路径
        :param img_savepath: 切分后图片保存目录,不传则自动生成
        :param part_height: 每段预设高度,默认1000像素
        :return: 切图保存目录、图片后缀
        '''
        # 获取图片后缀名(.jpg/.png等)
        img_text =os.path.splitext(img_path)[1] 
        
        # 判断图片文件是否存在
        if os.path.exists(img_path) :
            # 打开图片
            img = Image.open(img_path)
            
            # 检查图像是否为RGBA模式(带透明通道)
            if img.mode == 'RGBA':
                # 将RGBA模式转换为RGB模式(OCR不支持透明通道)
                img = img.convert('RGB')
                
            # 获取图片的宽度和高度
            width, height = img.size
            
            # 如果未传入保存路径,则自动生成:与原图同目录,以原文件名命名的文件夹
            if not img_savepath:
                filepath,filename=os.path.split(img_path)
                img_savepath=os.path.join(filepath,os.path.splitext(filename)[0])
            
            # 如果保存目录不存在,创建目录
            if not os.path.exists(img_savepath):
                os.makedirs(img_savepath)
            
            # 打印图片总高度和保存路径
            print ("图片的总高:",height,img_savepath)
            
            # 初始化切图参数
            i=0                 # 切图序号
            height_head=0       # 每一段的起始高度
            _height=part_height # 每一段的结束高度(初始值)
            
            # 循环切图,直到切完整个图片高度
            while  _height <= height:
                # 内层循环:智能寻找无文字区域作为切分点
                while 1:
                    '''
                    判断截图位置有没有文字,判断逻辑:
                    1. 截取当前切分点附近一小条图片(宽度=全图,高度=2像素)
                    2. 二值化处理
                    3. 查找图像中的连通组件(有相同特征的像素区域)
                    4. 连通组件数量<=3 → 判定为无文字区域,可以切分
                    '''
                    # 截取切分点上方2个像素的窄条,用于判断是否有文字
                    part=img.crop((0, _height-2, width,_height))
                    
                    # 将PIL图片转为numpy数组,供OpenCV处理
                    chu_img_array = np.asarray(part)
                    
                    # 判断是否为灰度图,不是则转为灰度图
                    if chu_img_array.ndim == 2:  # 已经是灰度图
                        gray = chu_img_array
                    else:  # 彩色图才进行转换
                        gray = cv2.cvtColor(chu_img_array, cv2.COLOR_BGR2GRAY)
                    
                    # 二值化处理(黑底白字)
                    thresh = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)[1]
                    
                    # 计算连通组件数量
                    num_labels, labels = cv2.connectedComponents(thresh)
                    
                    # 连通组件数量≤3 → 无文字 → 退出寻找循环
                    if num_labels <=3:
                        break
                    else:
                        # 有文字,向下移动4像素继续寻找
                        _height += 4 
                        
                        # 防止越界,到达图片底部直接停止
                        if _height >= height:
                            _height =height
                            break
                        
                        # 最多向下寻找20次(80像素),避免死循环
                        if _height >= part_height + 4*20:
                            break
                
                # 找到无文字区域,执行截图
                part = img.crop((0, height_head, width,_height))
                # 拼接切图保存路径
                _save_path = os.path.join(img_savepath, f"{i}{img_text}")
                # 保存切分后的小图
                part.save(_save_path)
                
                # 更新参数,准备下一次切图
                i +=1                # 切图序号+1
                height_head = _height # 下一段起始高度 = 当前结束高度
                
                # 已切到图片底部,退出循环
                if _height ==height:
                    break
                
                # 更新下一段结束高度
                _height +=part_height
                # 防止超出图片高度
                _height =_height  if _height < height else height
            
            # 返回切图保存目录和图片后缀
            return img_savepath,img_text
        else:
            # 文件不存在,抛出异常
            raise Exception("file no exists")

    # 判断图的大小执行不同的处理:长图走切分逻辑,短图直接识别
    def ocr_txt_new(self,img_path,width = 100, height = 100): 
        """
        对外主接口:OCR识别入口函数,自动判断图片类型并处理
        :param img_path: 图片路径
        :param width: 图片宽度
        :param height: 图片高度
        :return: 识别后的纯文本
        """
        # 定义缓存文件路径:图片同名+.txt
        _cache_img_content = f"{img_path}.txt"
        
        # 如果缓存文件已存在,直接读取返回,避免重复识别
        if os.path.exists(_cache_img_content):
            with open(_cache_img_content,"r",encoding="utf-8") as f:
                return f.read()
        
        # 特殊处理GIF:OCR不支持GIF,使用最后一帧转为PNG
        _ocr_path = img_path
        if img_path.endswith(".gif"): 
            _tmp_ocr_file = img_path[:-4]+"_gif.png"
            _ok = convert_gif_png(img_path,_tmp_ocr_file)
            if _ok:
                _ocr_path = _tmp_ocr_file
        
        # 根据图片高度选择识别方式:>1010像素 → 长图切分识别,否则→单图识别
        _content = ""
        if height > 1010 and not _ocr_path.endswith(".gif"):
            _content = self.ocr_txt_1(_ocr_path)  # 长图+切图识别
        else:
            _content =  self.ocr_txt_one(_ocr_path) # 短图直接识别
        
        # 将识别结果写入缓存文件,下次直接读取
        with open(_cache_img_content,"w",encoding="utf-8") as f:
            f.write(_content)
        
        return _content

    # 处理单个图片:直接识别,不切图
    def ocr_txt_one(self,img_path):
        '''
        单张图片直接OCR识别
        :param img_path: 图片路径
        :return: 按行拼接的纯文本
        '''
        # 执行OCR识别,cls=True开启方向分类
        result = self.OCR.ocr(img_path, cls=True)
        
        # 无识别结果,返回空字符串
        if not result:
            return ""
        
        _textlist = []
        # 遍历OCR返回结果
        for res in result:   
            if not res:
                continue
            for line in res:
                # line[0]:文字坐标四点
                # line[1][0]:识别出的文字内容
                txt=line[1][0]
                _textlist.append(txt)
        
        # 用换行符拼接所有文字
        return "\n".join(_textlist)

    def ocr_txt_1(self,img_path):
        """
        长图识别:先调用img_cut切图 → 逐张识别 → 坐标合并 → 按行排序 → 输出完整文本
        最后自动清理切图临时文件
        """
        # 获取文件后缀
        filemain,_text=os.path.splitext(img_path)
        print(f"imgcut:{img_path},{_text}")
        
        # 调用智能切图函数,返回切图目录和后缀
        filedir,_text=self.img_cut(img_path)
        
        # 获取切图目录下所有文件,并按创建时间排序
        filelist=list(os.listdir(filedir))
        filelist.sort()
        dir_list = sorted(filelist, key=lambda x: os.path.getctime(os.path.join(filedir, x)))
        filelist=[os.path.join(filedir,_f) for _f in dir_list]
        
        # 如果切图为空,直接识别原图
        if len(filelist)==0:
            filelist=[img_path]
        
        print("imgcut列表:",filelist)
        
        # 初始化变量:累计高度、坐标列表、文字列表
        height=0
        coordinatelist=[]
        txts=[]
        
        # 遍历所有切分后的小图,逐张OCR识别
        for ind,imgfile in enumerate(filelist):
            result = self.OCR.ocr(imgfile, cls=True)
            if not result:
                continue
            
            # 遍历识别结果
            for res in result:
                if not res:
                    continue
                for line in res:
                    coordinate=line[0]  # 文字四点坐标
                    txt=line[1][0]      # 识别文字
                    
                    # 拼接坐标:如果是第2张及以后的图,坐标Y轴需要加上前面图片的高度
                    if height:
                        coordinatelist.append([[i[0], i[1] + height] for i in coordinate])
                    else:
                        coordinatelist.append(coordinate)
                    
                    # 保存文字,带上切图序号
                    txts.append(str(ind)+"###"+txt)
            
            # 更新累计高度:以上一行文字底部+16像素作为下一段的基准高度
            if len(coordinatelist) > 0:
                height +=(coordinatelist[-1][2][1]+16)
            else:
                height += 1000 + 16
        
        # 无任何识别内容,返回空
        if len(coordinatelist)==0 and len(txts)==0:
            return ""
        
        # 将坐标列表转为numpy数组,用于计算行高
        arr=np.array(coordinatelist)
        # 提取所有坐标的Y轴数值
        arr_2d = arr[:, :, 1]
        # 计算最小行高,用于判断是否为同一行
        row_height= min(arr_2d[:,3]-arr_2d[:,0])
        
        # 按坐标Y(行)→ X(列)排序,保证文字从上到下、从左到右
        sortcoordinatelist=sorted(coordinatelist,key=lambda x:(x[0][1],x[0][0]))
        
        textlist=[]
        ts=[]
        lastheight=0
        
        # 按行分组文字:Y坐标差小于最小行高 → 同一行
        for t in sortcoordinatelist:
            _txt=txts[coordinatelist.index(t)]
            if lastheight==0:
                ts.append(((t[0][0],t[2][0]),_txt)) 
                lastheight=t[0][1] 
            elif t[0][1]-lastheight > row_height:
                # 换行
                ts.sort(key=lambda x:x[0][0])
                textlist.append(ts)
                ts=[((t[0][0],t[2][0]),_txt)]
                lastheight = t[0][1]
            else:
                # 同一行,追加
                ts.append(((t[0][0],t[2][0]),_txt))
        
        textlist.append(ts)        
        
        # 拼接最终文本
        textall = ""
        for ind,l in enumerate(textlist):
            text = "\t".join([t[1] for t in l])
            k = text.split("###")[0]
            v = text.replace(str(k)+"###","")
            textall += v+"\n"
        
        # 清理临时切图文件和目录
        if len(filelist) > 1:
            filedir = os.path.dirname(filelist[0])
            if os.path.exists(filedir):
                # 删除所有切图
                for f in filelist:
                    os.remove(f)
                # 删除切图目录
                os.rmdir(filedir)
        
        return textall


if __name__ == '__main__':
    # 测试代码:传入图片路径,执行OCR识别并打印结果
    img_path = r'/Users/ziguangchu/Downloads/tbc.jpeg'
    pd = Paddle_OCR()
    print (pd.ocr_txt_new(img_path,1080,3043))

项目分区导航hwcloud_api ⬅️ | 05-ocr_api | ➡️ openai4o_api