Merge branch 'github-dev'

This commit is contained in:
dengwenjie
2025-08-06 09:44:47 +08:00
4 changed files with 96 additions and 41 deletions

View File

@@ -4,6 +4,7 @@ import lombok.Data;
/**
* OCR 识别配置
*
* @author dwj
*/
@Data

View File

@@ -20,6 +20,7 @@ public class OcrInfo {
private String fullText;
private String base64Img;
public OcrInfo(List<List<OcrItem>> lineList, String fullText) {

View File

@@ -98,7 +98,23 @@ public interface OcrCommonRecModel extends AutoCloseable{
default BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize, OcrRecOptions options){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 识别并绘制Base64结果
* @param imageData 图片字节数组
* @return
*/
default String recognizeAndDrawToBase64(byte[] imageData, int fontSize, OcrRecOptions options){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 识别并绘制结果
* @param imageData 图片字节数组
* @return
*/
default OcrInfo recognizeAndDraw(byte[] imageData, int fontSize, OcrRecOptions options){
throw new UnsupportedOperationException("默认不支持该功能");
}
default List<OcrInfo> batchRecognize(List<BufferedImage> imageList, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");

View File

@@ -1,6 +1,5 @@
package cn.smartjavaai.ocr.model.common.recognize;
import ai.djl.Device;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
@@ -12,7 +11,7 @@ import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ModelZoo;
import ai.djl.repository.zoo.ZooModel;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.hutool.core.img.ImgUtil;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils;
@@ -28,7 +27,6 @@ import cn.smartjavaai.ocr.utils.OcrUtils;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.ObjectPool;
import org.apache.commons.pool2.impl.GenericObjectPool;
import org.opencv.core.Mat;
@@ -43,6 +41,7 @@ import java.util.stream.Collectors;
/**
* PPOCRV5 识别模型
*
* @author dwj
*/
@Slf4j
@@ -59,8 +58,8 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
private OcrCommonDetModel textDetModel;
@Override
public void loadModel(OcrRecModelConfig config){
if(StringUtils.isBlank(config.getRecModelPath())){
public void loadModel(OcrRecModelConfig config) {
if (StringUtils.isBlank(config.getRecModelPath())) {
throw new OcrException("recModelPath is null");
}
this.config = config;
@@ -68,11 +67,11 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
this.textDetModel = config.getTextDetModel();
//初始化 识别Criteria
Criteria<Image, String> recCriteria = OcrCommonRecCriterialFactory.createCriteria(config);
try{
try {
recognitionModel = ModelZoo.loadModel(recCriteria);
this.recPredictorPool = new GenericObjectPool<>(new PredictorFactory<>(recognitionModel));
int predictorPoolSize = config.getPredictorPoolSize();
if(config.getPredictorPoolSize() <= 0){
if (config.getPredictorPoolSize() <= 0) {
predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数
}
recPredictorPool.setMaxTotal(predictorPoolSize);
@@ -88,10 +87,10 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
@Override
public OcrInfo recognize(String imagePath, OcrRecOptions options) {
if(StringUtils.isBlank(config.getRecModelPath())){
if (StringUtils.isBlank(config.getRecModelPath())) {
throw new OcrException("recModelPath为空无法识别");
}
if(!FileUtils.isFileExists(imagePath)){
if (!FileUtils.isFileExists(imagePath)) {
throw new OcrException("图像文件不存在");
}
Image img = null;
@@ -101,14 +100,13 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
} catch (IOException e) {
throw new OcrException("无效的图片", e);
} finally {
if(img != null){
((Mat)img.getWrappedImage()).release();
if (img != null) {
((Mat) img.getWrappedImage()).release();
}
}
}
/**
*
* @param image
* @param options
* @return
@@ -116,7 +114,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
@Override
public OcrInfo recognize(Image image, OcrRecOptions options) {
List<OcrInfo> result = batchRecognizeDJLImage(Collections.singletonList(image), options);
if(CollectionUtils.isEmpty(result)){
if (CollectionUtils.isEmpty(result)) {
throw new OcrException("OCR识别结果为空");
}
return result.get(0);
@@ -125,12 +123,13 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
/**
* 批量矫正文本框
*
* @param boxList
* @param srcMat
* @param manager
* @return
*/
private List<Image> batchAlign(List<OcrBox> boxList, Mat srcMat,NDManager manager){
private List<Image> batchAlign(List<OcrBox> boxList, Mat srcMat, NDManager manager) {
List<Image> imageList = new ArrayList<>(boxList.size());
for (int i = 0; i < boxList.size(); i++) {
//透视变换 + 裁剪
@@ -149,12 +148,13 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
/**
* 批量矫正文本框
*
* @param itemList
* @param srcMat
* @param manager
* @return
*/
private List<Image> batchAlignWithDirection(List<OcrItem> itemList, Mat srcMat,NDManager manager){
private List<Image> batchAlignWithDirection(List<OcrItem> itemList, Mat srcMat, NDManager manager) {
List<Image> imageList = new ArrayList<>(itemList.size());
for (OcrItem ocrItem : itemList) {
//放射变换+裁剪
@@ -168,7 +168,6 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
}
// private RotatedBox recognize(OcrBox box,Mat srcMat,Predictor<Image, String> recPredictor,NDManager manager){
// try {
// //透视变换 + 裁剪
@@ -192,11 +191,12 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
/**
* 后处理:排序,分行
*
* @param rotatedBoxes
*/
private OcrInfo postProcessOcrResult(List<RotatedBox> rotatedBoxes, OcrRecOptions ocrRecOptions){
private OcrInfo postProcessOcrResult(List<RotatedBox> rotatedBoxes, OcrRecOptions ocrRecOptions) {
//不分行
if(!ocrRecOptions.isEnableLineSplit()){
if (!ocrRecOptions.isEnableLineSplit()) {
return OcrUtils.convertRotatedBoxesToOcrItems(rotatedBoxes);
}
//Y坐标升序排序
@@ -233,13 +233,13 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
@Override
public void recognizeAndDraw(String imagePath, String outputPath, int fontSize, OcrRecOptions options) {
if(!FileUtils.isFileExists(imagePath)){
if (!FileUtils.isFileExists(imagePath)) {
throw new OcrException("图像文件不存在");
}
try {
Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
OcrInfo ocrInfo = recognize(img, options);
if(Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()){
if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) {
throw new OcrException("未检测到文字");
}
Mat wrappedImage = (Mat) img.getWrappedImage();
@@ -254,18 +254,18 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
@Override
public OcrInfo recognize(BufferedImage image, OcrRecOptions options) {
if(!ImageUtils.isImageValid(image)){
if (!ImageUtils.isImageValid(image)) {
throw new OcrException("图像无效");
}
Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
OcrInfo ocrInfo = recognize(img, options);
((Mat)img.getWrappedImage()).release();
((Mat) img.getWrappedImage()).release();
return ocrInfo;
}
@Override
public OcrInfo recognize(byte[] imageData, OcrRecOptions options) {
if(Objects.isNull(imageData)){
if (Objects.isNull(imageData)) {
throw new OcrException("图像无效");
}
try {
@@ -278,18 +278,55 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
@Override
public BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize, OcrRecOptions options) {
if(!ImageUtils.isImageValid(sourceImage)){
if (!ImageUtils.isImageValid(sourceImage)) {
throw new OcrException("图像无效");
}
Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage));
OcrInfo ocrInfo = recognize(img, options);
if(Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()){
if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) {
throw new OcrException("未检测到文字");
}
OcrUtils.drawRectWithText(sourceImage, ocrInfo, fontSize);
return sourceImage;
}
@Override
public String recognizeAndDrawToBase64(byte[] imageData, int fontSize, OcrRecOptions options) {
if (Objects.isNull(imageData)) {
throw new OcrException("图像无效");
}
OcrInfo ocrInfo = recognize(imageData, options);
if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) {
throw new OcrException("未检测到文字");
}
try {
BufferedImage sourceImage = ImageIO.read(new ByteArrayInputStream(imageData));
OcrUtils.drawRectWithText(sourceImage, ocrInfo, fontSize);
return ImgUtil.toBase64(sourceImage, "png");
} catch (IOException e) {
throw new OcrException("导出图片失败", e);
}
}
@Override
public OcrInfo recognizeAndDraw(byte[] imageData, int fontSize, OcrRecOptions options) {
if (Objects.isNull(imageData)) {
throw new OcrException("图像无效");
}
OcrInfo ocrInfo = recognize(imageData, options);
if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) {
throw new OcrException("未检测到文字");
}
try {
BufferedImage sourceImage = ImageIO.read(new ByteArrayInputStream(imageData));
OcrUtils.drawRectWithText(sourceImage, ocrInfo, fontSize);
ocrInfo.setBase64Img(ImgUtil.toBase64(sourceImage, "png"));
return ocrInfo;
} catch (IOException e) {
throw new OcrException("导出图片失败", e);
}
}
@Override
public List<OcrInfo> batchRecognize(List<BufferedImage> imageList, OcrRecOptions options) {
List<Image> djlImageList = new ArrayList<>(imageList.size());
@@ -301,25 +338,25 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
} catch (Exception e) {
throw new OcrException(e);
} finally {
djlImageList.forEach(image -> ((Mat)image.getWrappedImage()).release());
djlImageList.forEach(image -> ((Mat) image.getWrappedImage()).release());
}
}
@Override
public List<OcrInfo> batchRecognizeDJLImage(List<Image> imageList, OcrRecOptions options) {
if(Objects.isNull(textDetModel)){
if (Objects.isNull(textDetModel)) {
throw new OcrException("textDetModel is null");
}
OcrRecOptions ocrRecOptions = options;
if(Objects.isNull(options)){
if (Objects.isNull(options)) {
ocrRecOptions = new OcrRecOptions();
}
if(CollectionUtils.isEmpty(imageList)){
if (CollectionUtils.isEmpty(imageList)) {
throw new OcrException("imageList is empty");
}
//检测文本
List<List<OcrBox>> boxeList = textDetModel.batchDetectDJLImage(imageList);
if(CollectionUtils.isEmpty(boxeList) || boxeList.size() != imageList.size()){
if (CollectionUtils.isEmpty(boxeList) || boxeList.size() != imageList.size()) {
throw new OcrException("未检测到文本");
}
Predictor<Image, String> predictor = null;
@@ -328,15 +365,15 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
predictor = recPredictorPool.borrowObject();
List<Image> allImageAlignList = new ArrayList<Image>();
//检测方向
if(ocrRecOptions.isEnableDirectionCorrect()){
if(Objects.isNull(directionModel)){
if (ocrRecOptions.isEnableDirectionCorrect()) {
if (Objects.isNull(directionModel)) {
throw new OcrException("请配置方向模型");
}
List<Mat> matList = imageList.stream()
.map(image -> (Mat)image.getWrappedImage())
.map(image -> (Mat) image.getWrappedImage())
.collect(Collectors.toList());
List<List<OcrItem>> ocrItemList = directionModel.batchDetect(boxeList, matList);
if(CollectionUtils.isEmpty(ocrItemList) || ocrItemList.size() != imageList.size()){
if (CollectionUtils.isEmpty(ocrItemList) || ocrItemList.size() != imageList.size()) {
throw new OcrException("方向检测失败");
}
allImageAlignList = new ArrayList<Image>();
@@ -348,7 +385,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
// }
allImageAlignList.addAll(imageAlignList);
}
}else{
} else {
for (int i = 0; i < boxeList.size(); i++) {
Mat srcMat = (Mat) imageList.get(i).getWrappedImage();
List<Image> imageAlignList = batchAlign(boxeList.get(i), srcMat, manager);
@@ -362,8 +399,8 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
int textIndex = 0;
for (int i = 0; i < boxeList.size(); i++) {
List<RotatedBox> rotatedBoxes = new ArrayList<>();
for (int j = 0; j < boxeList.get(i).size(); j++){
if(textIndex >= textList.size()){
for (int j = 0; j < boxeList.get(i).size(); j++) {
if (textIndex >= textList.size()) {
throw new OcrException("识别失败: 第" + i + "张图片, 第" + j + "个文本块,未识别到文本");
}
OcrBox box = boxeList.get(i).get(j);
@@ -377,7 +414,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
return ocrInfoList;
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
} finally {
if (predictor != null) {
try {
recPredictorPool.returnObject(predictor); //归还
@@ -393,16 +430,16 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel {
}
}
private List<String> batchRecognize(List<Image> imageAlignList){
private List<String> batchRecognize(List<Image> imageAlignList) {
Predictor<Image, String> predictor = null;
try {
predictor = recPredictorPool.borrowObject();
List<String> textList = predictor.batchPredict(imageAlignList);
imageAlignList.forEach(subImg -> ((Mat)subImg.getWrappedImage()).release());
imageAlignList.forEach(subImg -> ((Mat) subImg.getWrappedImage()).release());
return textList;
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
} finally {
if (predictor != null) {
try {
recPredictorPool.returnObject(predictor); //归还