1、OCR:新增表格识别模型

2、OCR:新增9个通用模型
3、OCR:支持批量检测识别
4、OCR:新增更多参数,使用更加灵活
5、人脸识别:支持ID查询及分页获取人脸信息
6、活体检测:视频检测支持设置最大帧数
This commit is contained in:
dengwenjie
2025-07-18 12:28:06 +08:00
parent 5d1f074de5
commit 6fbca62e5b
77 changed files with 3544 additions and 842 deletions

View File

@@ -6,7 +6,7 @@
<parent>
<groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-parent</artifactId>
<version>1.0.19</version>
<version>1.0.20</version>
</parent>
<artifactId>smartjavaai-ocr</artifactId>
@@ -18,9 +18,31 @@
<artifactId>smartjavaai-common</artifactId>
<version>${project.version}</version>
</dependency>
<!-- <dependency>-->
<!-- <groupId>ai.djl.paddlepaddle</groupId>-->
<!-- <artifactId>paddlepaddle-engine</artifactId>-->
<!-- <version>0.22.1</version>-->
<!-- </dependency>-->
<!-- <dependency>-->
<!-- <groupId>ai.djl.paddlepaddle</groupId>-->
<!-- <artifactId>paddlepaddle-model-zoo</artifactId>-->
<!-- <version>0.22.1</version>-->
<!-- </dependency>-->
<dependency>
<groupId>org.apache.poi</groupId>
<artifactId>poi</artifactId>
<version>4.0.0</version>
</dependency>
<dependency>
<groupId>dom4j</groupId>
<artifactId>dom4j</artifactId>
<version>1.6.1</version>
</dependency>
</dependencies>
<version>1.0.19</version>
<version>1.0.20</version>
<name>smartjavaai-ocr</name>
<description>SmartJavaAI</description>
<url>https://github.com/geekwenjie/SmartJavaAI</url>

View File

@@ -1,8 +1,10 @@
package cn.smartjavaai.ocr.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import cn.smartjavaai.ocr.enums.DirectionModelEnum;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.recognize.OcrCommonRecModel;
import lombok.Data;
@@ -12,33 +14,22 @@ import lombok.Data;
* @date 2025/4/22
*/
@Data
public class DirectionModelConfig {
public class DirectionModelConfig extends ModelConfig {
/**
* 模型
*/
private DirectionModelEnum modelEnum;
/**
* 设备类型
*/
private DeviceEnum device;
/**
* 检测模型路径
*/
private String modelPath;
/**
* 检测模型
* 文本检测模型
*/
private CommonDetModelEnum detModelEnum;
/**
* 检测模型路径
*/
private String detModelPath;
private OcrCommonDetModel textDetModel;

View File

@@ -1,5 +1,6 @@
package cn.smartjavaai.ocr.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import lombok.Data;
@@ -10,18 +11,13 @@ import lombok.Data;
* @date 2025/4/22
*/
@Data
public class OcrDetModelConfig {
public class OcrDetModelConfig extends ModelConfig {
/**
* 模型
*/
private CommonDetModelEnum modelEnum;
/**
* 设备类型
*/
private DeviceEnum device;
/**
* 检测模型路径
*/

View File

@@ -1,9 +1,12 @@
package cn.smartjavaai.ocr.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import cn.smartjavaai.ocr.enums.CommonRecModelEnum;
import cn.smartjavaai.ocr.enums.DirectionModelEnum;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.direction.OcrDirectionModel;
import lombok.Data;
/**
@@ -12,41 +15,26 @@ import lombok.Data;
* @date 2025/4/22
*/
@Data
public class OcrRecModelConfig {
/**
* 检测模型
*/
private CommonDetModelEnum detModelEnum;
public class OcrRecModelConfig extends ModelConfig {
/**
* 识别模型
*/
private CommonRecModelEnum recModelEnum;
/**
* 设备类型
*/
private DeviceEnum device;
/**
* 检测模型路径
*/
private String detModelPath;
/**
* 识别模型路径
*/
private String recModelPath;
/**
* 方向检测模型
* 文本检测模型
*/
private DirectionModelEnum directionModelEnum;
private OcrCommonDetModel textDetModel;
/**
* 方向检测模型路径
* 文本方向模型
*/
private String directionModelPath;
private OcrDirectionModel directionModel;
}

View File

@@ -0,0 +1,30 @@
package cn.smartjavaai.ocr.config;
import lombok.Data;
/**
* OCR 识别配置
* @author dwj
*/
@Data
public class OcrRecOptions {
/**
* 是否进行文本方向矫正
*/
private boolean enableDirectionCorrect = false;
/**
* 是否进行结果分行
*/
private boolean enableLineSplit = true;
public OcrRecOptions(boolean enableDirectionCorrect, boolean enableLineSplit) {
this.enableDirectionCorrect = enableDirectionCorrect;
this.enableLineSplit = enableLineSplit;
}
public OcrRecOptions() {
}
}

View File

@@ -0,0 +1,27 @@
package cn.smartjavaai.ocr.config;
import cn.smartjavaai.common.config.ModelConfig;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import cn.smartjavaai.ocr.enums.TableStructureModelEnum;
import lombok.Data;
/**
* OCR表格结构识别模型配置
* @author dwj
*/
@Data
public class TableStructureConfig extends ModelConfig {
/**
* 模型
*/
private TableStructureModelEnum modelEnum;
/**
* 检测模型路径
*/
private String modelPath;
}

View File

@@ -1,5 +1,6 @@
package cn.smartjavaai.ocr.entity;
import cn.smartjavaai.common.entity.DetectionRectangle;
import cn.smartjavaai.common.entity.Point;
import lombok.Data;
@@ -49,4 +50,23 @@ public class OcrBox {
(float)bottomLeft.getX(), (float)bottomLeft.getY()
};
}
/**
* 转换为 DetectionRectangle使用最小外包矩形
*/
public DetectionRectangle toDetectionRectangle() {
float[] pts = toFloatArray();
float minX = Math.min(Math.min(pts[0], pts[2]), Math.min(pts[4], pts[6]));
float minY = Math.min(Math.min(pts[1], pts[3]), Math.min(pts[5], pts[7]));
float maxX = Math.max(Math.max(pts[0], pts[2]), Math.max(pts[4], pts[6]));
float maxY = Math.max(Math.max(pts[1], pts[3]), Math.max(pts[5], pts[7]));
DetectionRectangle rect = new DetectionRectangle();
rect.setX((int) minX);
rect.setY((int) minY);
rect.setWidth((int) (maxX - minX));
rect.setHeight((int) (maxY - minY));
return rect;
}
}

View File

@@ -4,6 +4,7 @@ import lombok.Data;
import java.util.ArrayList;
import java.util.List;
import java.util.stream.Collectors;
/**
* OCR信息
@@ -15,13 +16,22 @@ public class OcrInfo {
private List<List<OcrItem>> lineList;
private List<OcrItem> ocrItemList;
private String fullText;
public OcrInfo(List<List<OcrItem>> lineList, String fullText) {
this.lineList = lineList;
this.fullText = fullText;
}
public OcrInfo() {
}
public List<OcrItem> flattenLines() {
return lineList.stream()
.flatMap(List::stream)
.collect(Collectors.toList());
}
}

View File

@@ -31,7 +31,6 @@ public class OcrItem {
private float score;
public OcrItem(OcrBox ocrBox, String text) {
this.ocrBox = ocrBox;
this.text = text;

View File

@@ -0,0 +1,33 @@
package cn.smartjavaai.ocr.entity;
import lombok.Data;
import java.util.List;
/**
* @author dwj
*/
@Data
public class TableStructureResult {
private List<OcrItem> ocrItemList;
private List<String> tableTagList;
private String html;
public TableStructureResult(List<OcrItem> ocrItemList, List<String> tableTagList) {
this.ocrItemList = ocrItemList;
this.tableTagList = tableTagList;
}
public TableStructureResult() {
}
public TableStructureResult(List<OcrItem> ocrItemList, List<String> tableTagList, String html) {
this.ocrItemList = ocrItemList;
this.tableTagList = tableTagList;
this.html = html;
}
}

View File

@@ -3,11 +3,16 @@ package cn.smartjavaai.ocr.enums;
/**
* OCR检测模型枚举
* @author dwj
* @date 2025/4/4
*/
public enum CommonDetModelEnum {
PADDLEOCR_V5_DET_MODEL;
PP_OCR_V5_SERVER_DET_MODEL,
PP_OCR_V5_MOBILE_DET_MODEL,
PP_OCR_V4_SERVER_DET_MODEL,
PP_OCR_V4_MOBILE_DET_MODEL;
/**

View File

@@ -3,11 +3,16 @@ package cn.smartjavaai.ocr.enums;
/**
* OCR识别模型枚举
* @author dwj
* @date 2025/4/4
*/
public enum CommonRecModelEnum {
PADDLEOCR_V5_REC_MODEL;
PP_OCR_V5_SERVER_REC_MODEL,
PP_OCR_V5_MOBILE_REC_MODEL,
PP_OCR_V4_SERVER_REC_MODEL,
PP_OCR_V4_MOBILE_REC_MODEL;
/**

View File

@@ -7,7 +7,12 @@ package cn.smartjavaai.ocr.enums;
*/
public enum DirectionModelEnum {
CH_PPOCR_MOBILE_V2_CLS;
CH_PPOCR_MOBILE_V2_CLS,
PP_LCNET_X0_25,
PP_LCNET_X1_0;
/**

View File

@@ -0,0 +1,30 @@
package cn.smartjavaai.ocr.enums;
/**
* OCR表格结构模型枚举
* @author dwj
*/
public enum TableStructureModelEnum {
SLANET,
//SLANEXT_WIRED,
SLANET_PLUS;
/**
* 根据名称获取枚举 (忽略大小写和下划线变体)
*/
public static TableStructureModelEnum fromName(String name) {
String formatted = name.trim().toUpperCase().replaceAll("[-_]", "");
for (TableStructureModelEnum model : values()) {
if (model.name().replaceAll("_", "").equals(formatted)) {
return model;
}
}
throw new IllegalArgumentException("未知模型名称: " + name);
}
}

View File

@@ -3,14 +3,17 @@ package cn.smartjavaai.ocr.factory;
import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.ocr.config.DirectionModelConfig;
import cn.smartjavaai.ocr.config.OcrRecModelConfig;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import cn.smartjavaai.ocr.enums.CommonRecModelEnum;
import cn.smartjavaai.ocr.enums.DirectionModelEnum;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.detect.PpOCRV5DetModel;
import cn.smartjavaai.ocr.model.common.recognize.OcrCommonRecModel;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModelImpl;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.model.common.direction.OcrDirectionModel;
import cn.smartjavaai.ocr.model.common.direction.PPOCRMobileV2Model;
import cn.smartjavaai.ocr.model.common.recognize.PpOCRV5RecModel;
import cn.smartjavaai.ocr.model.common.direction.PPOCRMobileV2ClsModel;
import cn.smartjavaai.ocr.model.common.recognize.OcrCommonRecModel;
import cn.smartjavaai.ocr.model.common.recognize.OcrCommonRecModelImpl;
import lombok.extern.slf4j.Slf4j;
import java.util.Map;
@@ -27,29 +30,29 @@ public class OcrModelFactory {
// 使用 volatile 和双重检查锁定来确保线程安全的单例模式
private static volatile OcrModelFactory instance;
private static final ConcurrentHashMap<String, OcrCommonDetModel> commonDetModelMap = new ConcurrentHashMap<>();
private static final ConcurrentHashMap<CommonDetModelEnum, OcrCommonDetModel> commonDetModelMap = new ConcurrentHashMap<>();
private static final ConcurrentHashMap<String, OcrCommonRecModel> commonRecModelMap = new ConcurrentHashMap<>();
private static final ConcurrentHashMap<CommonRecModelEnum, OcrCommonRecModel> commonRecModelMap = new ConcurrentHashMap<>();
private static final ConcurrentHashMap<String, OcrDirectionModel> directionModelMap = new ConcurrentHashMap<>();
private static final ConcurrentHashMap<DirectionModelEnum, OcrDirectionModel> directionModelMap = new ConcurrentHashMap<>();
/**
* 检测模型注册表
*/
private static final Map<String, Class<? extends OcrCommonDetModel>> commonDetRegistry =
private static final Map<CommonDetModelEnum, Class<? extends OcrCommonDetModel>> commonDetRegistry =
new ConcurrentHashMap<>();
/**
* 识别模型注册表
*/
private static final Map<String, Class<? extends OcrCommonRecModel>> commonRecRegistry =
private static final Map<CommonRecModelEnum, Class<? extends OcrCommonRecModel>> commonRecRegistry =
new ConcurrentHashMap<>();
/**
* 方向分类模型注册表
*/
private static final Map<String, Class<? extends OcrDirectionModel>> directionRegistry =
private static final Map<DirectionModelEnum, Class<? extends OcrDirectionModel>> directionRegistry =
new ConcurrentHashMap<>();
@@ -68,29 +71,29 @@ public class OcrModelFactory {
/**
* 注册通用检测模型
* @param name
* @param detModelEnum
* @param clazz
*/
private static void registerCommonDetModel(String name, Class<? extends OcrCommonDetModel> clazz) {
commonDetRegistry.put(name.toLowerCase(), clazz);
private static void registerCommonDetModel(CommonDetModelEnum detModelEnum, Class<? extends OcrCommonDetModel> clazz) {
commonDetRegistry.put(detModelEnum, clazz);
}
/**
* 注册通用识别模型
* @param name
* @param recModelEnum
* @param clazz
*/
private static void registerCommonRecModel(String name, Class<? extends OcrCommonRecModel> clazz) {
commonRecRegistry.put(name.toLowerCase(), clazz);
private static void registerCommonRecModel(CommonRecModelEnum recModelEnum, Class<? extends OcrCommonRecModel> clazz) {
commonRecRegistry.put(recModelEnum, clazz);
}
/**
* 注册通用方向分类模型
* @param name
* @param directionModelEnum
* @param clazz
*/
private static void registerDirectionModel(String name, Class<? extends OcrDirectionModel> clazz) {
directionRegistry.put(name.toLowerCase(), clazz);
private static void registerDirectionModel(DirectionModelEnum directionModelEnum, Class<? extends OcrDirectionModel> clazz) {
directionRegistry.put(directionModelEnum, clazz);
}
@@ -103,7 +106,7 @@ public class OcrModelFactory {
if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){
throw new OcrException("未配置OCR模型");
}
return commonDetModelMap.computeIfAbsent(config.getModelEnum().name(), k -> {
return commonDetModelMap.computeIfAbsent(config.getModelEnum(), k -> {
return createCommonDetModel(config);
});
}
@@ -117,7 +120,7 @@ public class OcrModelFactory {
if(Objects.isNull(config) || Objects.isNull(config.getRecModelEnum())){
throw new OcrException("未配置OCR模型");
}
return commonRecModelMap.computeIfAbsent(config.getRecModelEnum().name(), k -> {
return commonRecModelMap.computeIfAbsent(config.getRecModelEnum(), k -> {
return createCommonRecModel(config);
});
}
@@ -131,7 +134,7 @@ public class OcrModelFactory {
if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){
throw new OcrException("未配置OCR模型");
}
return directionModelMap.computeIfAbsent(config.getModelEnum().name(), k -> {
return directionModelMap.computeIfAbsent(config.getModelEnum(), k -> {
return createDirectionModel(config);
});
}
@@ -144,7 +147,7 @@ public class OcrModelFactory {
* @return
*/
private OcrCommonDetModel createCommonDetModel(OcrDetModelConfig config) {
Class<?> clazz = commonDetRegistry.get(config.getModelEnum().name().toLowerCase());
Class<?> clazz = commonDetRegistry.get(config.getModelEnum());
if(clazz == null){
throw new OcrException("Unsupported model");
}
@@ -165,7 +168,7 @@ public class OcrModelFactory {
* @return
*/
private OcrCommonRecModel createCommonRecModel(OcrRecModelConfig config) {
Class<?> clazz = commonRecRegistry.get(config.getRecModelEnum().name().toLowerCase());
Class<?> clazz = commonRecRegistry.get(config.getRecModelEnum());
if(clazz == null){
throw new OcrException("Unsupported model");
}
@@ -185,7 +188,7 @@ public class OcrModelFactory {
* @return
*/
private OcrDirectionModel createDirectionModel(DirectionModelConfig config) {
Class<?> clazz = directionRegistry.get(config.getModelEnum().name().toLowerCase());
Class<?> clazz = directionRegistry.get(config.getModelEnum());
if(clazz == null){
throw new OcrException("Unsupported model");
}
@@ -202,9 +205,18 @@ public class OcrModelFactory {
// 初始化默认算法
static {
registerCommonDetModel("PADDLEOCR_V5_DET_MODEL", PpOCRV5DetModel.class);
registerCommonRecModel("PADDLEOCR_V5_REC_MODEL", PpOCRV5RecModel.class);
registerDirectionModel("CH_PPOCR_MOBILE_V2_CLS", PPOCRMobileV2Model.class);
//通用-检测模型
registerCommonDetModel(CommonDetModelEnum.PP_OCR_V5_SERVER_DET_MODEL, OcrCommonDetModelImpl.class);
registerCommonDetModel(CommonDetModelEnum.PP_OCR_V5_MOBILE_DET_MODEL, OcrCommonDetModelImpl.class);
registerCommonDetModel(CommonDetModelEnum.PP_OCR_V4_SERVER_DET_MODEL, OcrCommonDetModelImpl.class);
registerCommonDetModel(CommonDetModelEnum.PP_OCR_V4_MOBILE_DET_MODEL, OcrCommonDetModelImpl.class);
registerCommonRecModel(CommonRecModelEnum.PP_OCR_V5_SERVER_REC_MODEL, OcrCommonRecModelImpl.class);
registerCommonRecModel(CommonRecModelEnum.PP_OCR_V5_MOBILE_REC_MODEL, OcrCommonRecModelImpl.class);
registerCommonRecModel(CommonRecModelEnum.PP_OCR_V4_SERVER_REC_MODEL, OcrCommonRecModelImpl.class);
registerCommonRecModel(CommonRecModelEnum.PP_OCR_V4_MOBILE_REC_MODEL, OcrCommonRecModelImpl.class);
registerDirectionModel(DirectionModelEnum.CH_PPOCR_MOBILE_V2_CLS, PPOCRMobileV2ClsModel.class);
registerDirectionModel(DirectionModelEnum.PP_LCNET_X0_25, PPOCRMobileV2ClsModel.class);
registerDirectionModel(DirectionModelEnum.PP_LCNET_X1_0, PPOCRMobileV2ClsModel.class);
log.debug("缓存目录:{}", Config.getCachePath());
}

View File

@@ -0,0 +1,119 @@
package cn.smartjavaai.ocr.factory;
import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.ocr.config.DirectionModelConfig;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.config.OcrRecModelConfig;
import cn.smartjavaai.ocr.config.TableStructureConfig;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import cn.smartjavaai.ocr.enums.CommonRecModelEnum;
import cn.smartjavaai.ocr.enums.DirectionModelEnum;
import cn.smartjavaai.ocr.enums.TableStructureModelEnum;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModelImpl;
import cn.smartjavaai.ocr.model.common.direction.OcrDirectionModel;
import cn.smartjavaai.ocr.model.common.direction.PPOCRMobileV2ClsModel;
import cn.smartjavaai.ocr.model.common.recognize.OcrCommonRecModel;
import cn.smartjavaai.ocr.model.common.recognize.OcrCommonRecModelImpl;
import cn.smartjavaai.ocr.model.table.CommonTableStructureModel;
import cn.smartjavaai.ocr.model.table.TableStructureModel;
import lombok.extern.slf4j.Slf4j;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* OCR 表格识别模型工厂
* @author dwj
*/
@Slf4j
public class TableRecModelFactory {
// 使用 volatile 和双重检查锁定来确保线程安全的单例模式
private static volatile TableRecModelFactory instance;
/**
* 模型缓存
*/
private static final ConcurrentHashMap<TableStructureModelEnum, TableStructureModel> tableStructureModelMap = new ConcurrentHashMap<>();
/**
* 模型注册表
*/
private static final Map<TableStructureModelEnum, Class<? extends TableStructureModel>> tableStructureRegistry =
new ConcurrentHashMap<>();
public static TableRecModelFactory getInstance() {
if (instance == null) {
synchronized (TableRecModelFactory.class) {
if (instance == null) {
instance = new TableRecModelFactory();
}
}
}
return instance;
}
/**
* 注册模型
* @param tableStructureModelEnum
* @param clazz
*/
private static void registerTableStructureModel(TableStructureModelEnum tableStructureModelEnum, Class<? extends TableStructureModel> clazz) {
tableStructureRegistry.put(tableStructureModelEnum, clazz);
}
/**
* 获取模型(通过配置)
* @param config
* @return
*/
public TableStructureModel getTableStructureModel(TableStructureConfig config) {
if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){
throw new OcrException("未配置OCR模型");
}
return tableStructureModelMap.computeIfAbsent(config.getModelEnum(), k -> {
return createTableStructureModel(config);
});
}
/**
* 创建模型
* @param config
* @return
*/
private TableStructureModel createTableStructureModel(TableStructureConfig config) {
Class<?> clazz = tableStructureRegistry.get(config.getModelEnum());
if(clazz == null){
throw new OcrException("Unsupported model");
}
TableStructureModel model = null;
try {
model = (TableStructureModel) clazz.newInstance();
} catch (InstantiationException | IllegalAccessException e) {
throw new OcrException(e);
}
model.loadModel(config);
return model;
}
// 初始化默认算法
static {
registerTableStructureModel(TableStructureModelEnum.SLANET, CommonTableStructureModel.class);
registerTableStructureModel(TableStructureModelEnum.SLANET_PLUS, CommonTableStructureModel.class);
log.debug("缓存目录:{}", Config.getCachePath());
}
}

View File

@@ -29,8 +29,6 @@ public interface OcrCommonDetModel extends AutoCloseable{
}
/**
* 文本检测
* @param image BufferedImage
@@ -77,4 +75,23 @@ public interface OcrCommonDetModel extends AutoCloseable{
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 文本检测(批量)
* @param imageList BufferedImage
* @return
*/
default List<List<OcrBox>> batchDetect(List<BufferedImage> imageList) {
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 文本检测(批量)
* @param imageList DJL Image
* @return
*/
default List<List<OcrBox>> batchDetectDJLImage(List<Image> imageList){
throw new UnsupportedOperationException("默认不支持该功能");
}
}

View File

@@ -1,20 +1,16 @@
package cn.smartjavaai.ocr.model.common.detect;
import ai.djl.Device;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
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 ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils;
@@ -22,7 +18,7 @@ import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.model.common.detect.translator.PPOCRV5DetTranslator;
import cn.smartjavaai.ocr.model.common.detect.criteria.OcrCommonDetCriterialFactory;
import cn.smartjavaai.ocr.utils.OcrUtils;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils;
@@ -38,18 +34,14 @@ import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.List;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
import java.util.*;
/**
* PPOCRV5 检测模型
* ocr通用检测模型实现类
* @author dwj
* @date 2025/4/21
*/
@Slf4j
public class PpOCRV5DetModel implements OcrCommonDetModel {
public class OcrCommonDetModelImpl implements OcrCommonDetModel{
private ObjectPool<Predictor<Image, NDList>> detPredictorPool;
@@ -62,22 +54,9 @@ public class PpOCRV5DetModel implements OcrCommonDetModel {
if(StringUtils.isBlank(config.getDetModelPath())){
throw new OcrException("modelPath is null");
}
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
}
this.config = config;
//初始化 检测Criteria
Criteria<Image, NDList> detCriteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, NDList.class)
.optModelPath(Paths.get(config.getDetModelPath()))
.optTranslator(new PPOCRV5DetTranslator(new ConcurrentHashMap<String, String>()))
.optDevice(device)
.optProgress(new ProgressBar())
.build();
Criteria<Image, NDList> detCriteria = OcrCommonDetCriterialFactory.createCriteria(config);
try{
detectionModel = ModelZoo.loadModel(detCriteria);
// 创建池子每个线程独享 Predictor
@@ -107,28 +86,9 @@ public class PpOCRV5DetModel implements OcrCommonDetModel {
@Override
public List<OcrBox> detect(Image image){
Predictor<Image, NDList> predictor = null;
try (NDManager manager = NDManager.newBaseManager()) {
predictor = detPredictorPool.borrowObject();
NDList result = predictor.predict(image);
result.attach(manager);
return OcrUtils.convertToOcrBox(result, image);
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
detPredictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
List<Image> imageList = Collections.singletonList(image);
List<List<OcrBox>> result = batchDetectDJLImage(imageList);
return result.get(0);
}
@Override
@@ -136,7 +96,7 @@ public class PpOCRV5DetModel implements OcrCommonDetModel {
if(!FileUtils.isFileExists(imagePath)){
throw new OcrException("图像文件不存在");
}
try (NDManager manager = NDManager.newBaseManager()) {
try {
Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
List<OcrBox> boxList = detect(img);
if(Objects.isNull(boxList) || boxList.isEmpty()){
@@ -194,10 +154,58 @@ public class PpOCRV5DetModel implements OcrCommonDetModel {
img.save(outputStream, "png");
// 将字节流转换为 BufferedImage
byte[] imageBytes = outputStream.toByteArray();
((Mat) img.getWrappedImage()).release();
return ImageIO.read(new ByteArrayInputStream(imageBytes));
} catch (IOException e) {
throw new OcrException("导出图片失败", e);
} finally {
if (img != null){
((Mat) img.getWrappedImage()).release();
}
}
}
@Override
public List<List<OcrBox>> batchDetect(List<BufferedImage> imageList) {
List<Image> djlImageList = new ArrayList<>(imageList.size());
try {
for (BufferedImage bufferedImage : imageList) {
djlImageList.add(ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(bufferedImage)));
}
return batchDetectDJLImage(djlImageList);
} catch (Exception e) {
throw new OcrException(e);
} finally {
djlImageList.forEach(image -> ((Mat)image.getWrappedImage()).release());
}
}
@Override
public List<List<OcrBox>> batchDetectDJLImage(List<Image> imageList) {
if(!ImageUtils.isAllImageSizeEqual(imageList)){
throw new OcrException("图片尺寸不一致");
}
Predictor<Image, NDList> predictor = null;
try (NDManager manager = NDManager.newBaseManager()) {
predictor = detPredictorPool.borrowObject();
List<NDList> result = predictor.batchPredict(imageList);
result.forEach(ndList -> ndList.attach(manager));
return OcrUtils.convertToOcrBox(result);
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
detPredictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@@ -218,4 +226,6 @@ public class PpOCRV5DetModel implements OcrCommonDetModel {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,53 @@
package cn.smartjavaai.ocr.model.common.detect.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.ndarray.NDList;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import cn.smartjavaai.ocr.model.common.detect.translator.PPOCRDetTranslator;
import org.apache.commons.lang3.StringUtils;
import java.nio.file.Paths;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* @author dwj
* @date 2025/7/8
*/
public class OcrCommonDetCriterialFactory {
public static Criteria<Image, NDList> createCriteria(OcrDetModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
}
Criteria<Image, NDList> criteria = null;
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
params.putAll(config.getCustomParams());
if(StringUtils.isNotBlank(config.getBatchifier())){
params.put("batchifier", config.getBatchifier());
}
if(config.getModelEnum() == CommonDetModelEnum.PP_OCR_V5_SERVER_DET_MODEL ||
config.getModelEnum() == CommonDetModelEnum.PP_OCR_V5_MOBILE_DET_MODEL ||
config.getModelEnum() == CommonDetModelEnum.PP_OCR_V4_SERVER_DET_MODEL ||
config.getModelEnum() == CommonDetModelEnum.PP_OCR_V4_MOBILE_DET_MODEL
){
criteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, NDList.class)
.optModelPath(Paths.get(config.getDetModelPath()))
.optTranslator(new PPOCRDetTranslator(params))
.optDevice(device)
.optProgress(new ProgressBar())
.build();
}
return criteria;
}
}

View File

@@ -27,24 +27,35 @@ import java.util.Map;
* @mail 179209347@qq.com
* @website www.aias.top
*/
public class PPOCRV5DetTranslator implements Translator<Image, NDList> {
public class PPOCRDetTranslator implements Translator<Image, NDList> {
// det_algorithm == "DB"
private final float thresh = 0.3f;
private final boolean use_dilation = false;
private final String score_mode = "fast";
private final String box_type = "quad";
//检测的图像边长限制
private final int limit_side_len;
//输出的最大文本框数量
private final int max_candidates;
//文本框最小尺寸阈值
private final int min_size;
//文本框的分数阈值
private final float box_thresh;
/**
* 这个参数是检测后处理时控制文本框大小的默认1.6可以尝试改成2.5或者更大反之如果觉得文本框不够紧凑也可以把该参数调小
* 检测框大小过于紧贴文字或检测框过大可以调整db_unclip_ratio这个参数加大参数可以扩大检测框减小参数可以减小检测框大小
*/
private final float unclip_ratio;
private float ratio_h;
private float ratio_w;
private int img_height;
private int img_width;
public PPOCRV5DetTranslator(Map<String, ?> arguments) {
private String batchifier;
public PPOCRDetTranslator(Map<String, ?> arguments) {
limit_side_len =
arguments.containsKey("limit_side_len")
? Integer.parseInt(arguments.get("limit_side_len").toString())
@@ -65,6 +76,10 @@ public class PPOCRV5DetTranslator implements Translator<Image, NDList> {
arguments.containsKey("unclip_ratio")
? Float.parseFloat(arguments.get("unclip_ratio").toString())
: 1.6f;
batchifier = arguments.containsKey("batchifier")
? arguments.get("batchifier").toString()
: "stack";
}
@Override
@@ -509,13 +524,13 @@ public class PPOCRV5DetTranslator implements Translator<Image, NDList> {
new float[]{0.485f, 0.456f, 0.406f},
new float[]{0.229f, 0.224f, 0.225f});
img = img.expandDims(0);
// img = img.expandDims(0);
return new NDList(img);
}
@Override
public Batchifier getBatchifier() {
return null;
return Batchifier.fromString(batchifier);
}
}

View File

@@ -8,6 +8,7 @@ import cn.smartjavaai.ocr.entity.DirectionInfo;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.entity.OcrInfo;
import cn.smartjavaai.ocr.entity.OcrItem;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import org.opencv.core.Mat;
import java.awt.image.BufferedImage;
@@ -19,6 +20,14 @@ import java.util.List;
*/
public interface OcrDirectionModel extends AutoCloseable{
default void setTextDetModel(OcrCommonDetModel detModel){
throw new UnsupportedOperationException("默认不支持该功能");
}
default OcrCommonDetModel getTextDetModel(){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 加载模型
* @param config
@@ -71,7 +80,11 @@ public interface OcrDirectionModel extends AutoCloseable{
* @param manager
* @return
*/
default List<OcrItem> detect(List<OcrBox> boxList, Mat srcMat, NDManager manager) {
default List<OcrItem> detect(List<OcrBox> boxList, Mat srcMat) {
throw new UnsupportedOperationException("默认不支持该功能");
}
default List<List<OcrItem>> batchDetect(List<List<OcrBox>> boxList, List<Mat> srcMatList) {
throw new UnsupportedOperationException("默认不支持该功能");
}

View File

@@ -6,9 +6,7 @@ import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.Point;
import ai.djl.ndarray.NDManager;
import ai.djl.opencv.OpenCVImageFactory;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ModelZoo;
@@ -20,17 +18,15 @@ import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.ocr.config.DirectionModelConfig;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.entity.*;
import cn.smartjavaai.ocr.enums.AngleEnum;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.factory.OcrModelFactory;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.direction.criteria.DirectionCriteriaFactory;
import cn.smartjavaai.ocr.model.common.direction.translator.PpWordRotateTranslator;
import cn.smartjavaai.ocr.opencv.OcrNDArrayUtils;
import cn.smartjavaai.ocr.opencv.OcrOpenCVUtils;
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;
@@ -45,9 +41,10 @@ import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Objects;
import java.util.UUID;
import java.util.concurrent.ConcurrentHashMap;
/**
* PPOCRMobileV2Model 方向分类模型
@@ -55,38 +52,34 @@ import java.util.UUID;
* @date 2025/4/21
*/
@Slf4j
public class PPOCRMobileV2Model implements OcrDirectionModel {
public class PPOCRMobileV2ClsModel implements OcrDirectionModel {
private ObjectPool<Predictor<Image, DirectionInfo>> predictorPool;
private DirectionModelConfig config;
private OcrCommonDetModel detModel;
private ZooModel<Image, DirectionInfo> model;
private OcrCommonDetModel textDetModel;
@Override
public void loadModel(DirectionModelConfig config){
if(StringUtils.isBlank(config.getModelPath())){
throw new OcrException("modelPath is null");
}
this.config = config;
this.textDetModel = config.getTextDetModel();
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
}
Criteria<Image, DirectionInfo> criteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, DirectionInfo.class)
.optModelPath(Paths.get(config.getModelPath()))
.optDevice(device)
.optTranslator(new PpWordRotateTranslator())
.optProgress(new ProgressBar())
.build();
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
if(StringUtils.isNotBlank(config.getBatchifier())){
params.put("batchifier", config.getBatchifier());
}
Criteria<Image, DirectionInfo> criteria = DirectionCriteriaFactory.createCriteria(config);
try{
model = ModelZoo.loadModel(criteria);
// 创建池子每个线程独享 Predictor
@@ -96,15 +89,6 @@ public class PPOCRMobileV2Model implements OcrDirectionModel {
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new OcrException("模型加载失败", e);
}
//获取检测模型
if(StringUtils.isNotBlank(config.getDetModelPath()) && Objects.nonNull(config.getDetModelEnum())){
OcrDetModelConfig detModelConfig = new OcrDetModelConfig();
detModelConfig.setModelEnum(config.getDetModelEnum());
detModelConfig.setDetModelPath(config.getDetModelPath());
detModelConfig.setDevice(config.getDevice());
detModel = OcrModelFactory.getInstance().getDetModel(detModelConfig);
}
}
@Override
@@ -115,124 +99,85 @@ public class PPOCRMobileV2Model implements OcrDirectionModel {
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
return detect(img);
} catch (IOException e) {
throw new OcrException("无效的图片", e);
}finally {
if(img != null){
((Mat)img.getWrappedImage()).release();
}
}
List<OcrItem> ocrItemList = detect(img);
((Mat)img.getWrappedImage()).release();
return ocrItemList;
}
@Override
public List<OcrItem> detect(Image image){
if(Objects.isNull(textDetModel)){
throw new OcrException("textDetModel is null");
}
//检测文本
List<OcrBox> boxeList = detModel.detect(image);
List<OcrBox> boxeList = textDetModel.detect(image);
if(Objects.isNull(boxeList) || boxeList.isEmpty()){
throw new OcrException("未检测到文本");
}
Predictor<Image, DirectionInfo> predictor = null;
List<OcrItem> ocrItemList = new ArrayList<>();
try (NDManager manager = NDManager.newBaseManager()) {
Mat srcMat = (Mat) image.getWrappedImage();
predictor = predictorPool.borrowObject();
for (OcrBox box : boxeList){
OcrItem ocrItem = detect(box, srcMat, predictor, manager);
ocrItemList.add(ocrItem);
}
return ocrItemList;
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
Mat srcMat = (Mat) image.getWrappedImage();
return detect(boxeList, srcMat);
}
/**
* 基于文本框检测方向
* @param box
* @param srcMat
* @param predictor
* @param manager
* @return
*/
private OcrItem detect(OcrBox box, Mat srcMat, Predictor<Image, DirectionInfo> predictor, NDManager manager){
if(Objects.isNull(box)){
throw new OcrException("box参数为空");
}
try {
//透视变换及裁剪
Image subImg = OcrUtils.transformAndCrop(srcMat, box);
DirectionInfo directionInfo = null;
String angle;
//高宽比 > 1.5 纵向
if (subImg.getHeight() * 1.0 / subImg.getWidth() > 1.5) {
//旋转图片90度
subImg = OcrUtils.rotateImg(manager, subImg);
//检测方向
directionInfo = predictor.predict(subImg);
if (directionInfo.getName().equalsIgnoreCase("Rotate")) {
angle = "270";
} else {
angle = "90";
}
}else{ //横向
directionInfo = predictor.predict(subImg);
if (directionInfo.getName().equalsIgnoreCase("No Rotate")) {
angle = "0";
} else {
angle = "180";
}
}
((Mat)subImg.getWrappedImage()).release();
return new OcrItem(box, AngleEnum.fromValue(angle), directionInfo.getProb().floatValue());
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}
}
// /**
// * 基于文本框检测方向
// * @param box
// * @param srcMat
// * @param predictor
// * @param manager
// * @return
// */
// private OcrItem detect(OcrBox box, Mat srcMat, Predictor<Image, DirectionInfo> predictor, NDManager manager){
// if(Objects.isNull(box)){
// throw new OcrException("box参数为空");
// }
// try {
// //透视变换及裁剪
// Image subImg = OcrUtils.transformAndCrop(srcMat, box);
// DirectionInfo directionInfo = null;
// String angle;
// //高宽比 > 1.5 纵向
// if (subImg.getHeight() * 1.0 / subImg.getWidth() > 1.5) {
// //旋转图片90度
// subImg = OcrUtils.rotateImg(manager, subImg);
// //检测方向
// directionInfo = predictor.predict(subImg);
// if (directionInfo.getName().equalsIgnoreCase("Rotate")) {
// angle = "270";
// } else {
// angle = "90";
// }
// }else{ //横向
// directionInfo = predictor.predict(subImg);
// if (directionInfo.getName().equalsIgnoreCase("No Rotate")) {
// angle = "0";
// } else {
// angle = "180";
// }
// }
// ((Mat)subImg.getWrappedImage()).release();
// return new OcrItem(box, AngleEnum.fromValue(angle), directionInfo.getProb().floatValue());
// } catch (Exception e) {
// throw new OcrException("OCR检测错误", e);
// }
// }
@Override
public List<OcrItem> detect(List<OcrBox> boxList,Mat srcMat,NDManager manager){
public List<OcrItem> detect(List<OcrBox> boxList,Mat srcMat){
if(Objects.isNull(boxList) || boxList.isEmpty()){
throw new OcrException("boxList为空");
}
Predictor<Image, DirectionInfo> predictor = null;
List<OcrItem> ocrItemList = new ArrayList<>();
try {
predictor = predictorPool.borrowObject();
for (OcrBox box : boxList){
OcrItem ocrItem = detect(box, srcMat, predictor, manager);
ocrItemList.add(ocrItem);
}
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
List<List<OcrItem>> ocrItemList = batchDetect(Collections.singletonList(boxList), Collections.singletonList(srcMat));
if(Objects.isNull(ocrItemList) || ocrItemList.isEmpty()){
throw new OcrException("方向检测失败");
}
return ocrItemList;
return ocrItemList.get(0);
}
@Override
@@ -240,8 +185,9 @@ public class PPOCRMobileV2Model implements OcrDirectionModel {
if(!FileUtils.isFileExists(imagePath)){
throw new OcrException("图像文件不存在");
}
try (NDManager manager = NDManager.newBaseManager()) {
Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
List<OcrItem> itemList = detect(img);
if(Objects.isNull(itemList) || itemList.isEmpty()){
throw new OcrException("未检测到文字");
@@ -250,9 +196,12 @@ public class PPOCRMobileV2Model implements OcrDirectionModel {
Path output = Paths.get(outputPath);
log.debug("Saving to {}", output.toAbsolutePath().toString());
img.save(Files.newOutputStream(output), "png");
((Mat) img.getWrappedImage()).release();
} catch (IOException e) {
throw new OcrException(e);
} finally {
if (img != null){
((Mat)img.getWrappedImage()).release();
}
}
}
@@ -297,13 +246,129 @@ public class PPOCRMobileV2Model implements OcrDirectionModel {
img.save(outputStream, "png");
// 将字节流转换为 BufferedImage
byte[] imageBytes = outputStream.toByteArray();
((Mat) img.getWrappedImage()).release();
return ImageIO.read(new ByteArrayInputStream(imageBytes));
} catch (IOException e) {
throw new OcrException("导出图片失败", e);
} finally {
if (img != null){
((Mat) img.getWrappedImage()).release();
}
}
}
@Override
public List<List<OcrItem>> batchDetect(List<List<OcrBox>> boxList, List<Mat> srcMatList) {
if(CollectionUtils.isEmpty(boxList)){
throw new OcrException("boxList 不能为空");
}
if(CollectionUtils.isEmpty(srcMatList)){
throw new OcrException("srcMatList 不能为空");
}
//检查参数
for (int i = 0; i < srcMatList.size(); i++) {
List<OcrBox> ocrBoxes = boxList.get(i);
Mat mat = srcMatList.get(i);
if (ocrBoxes == null) {
throw new OcrException("" + i + " 个 boxList 为 null");
}
if (ocrBoxes.isEmpty()) {
throw new OcrException("" + i + " 个 boxList 没有检测结果");
}
if (mat.empty()) {
throw new OcrException("" + i + " 张图片为空 Mat");
}
}
List<Image> imageList = new ArrayList<Image>();
List<Boolean> isRotatedList = new ArrayList<Boolean>();
int index = 0;
try (NDManager manager = model.getNDManager().newSubManager()){
for(int i = 0; i < srcMatList.size(); i++){
for (int j = 0; j < boxList.get(i).size(); j++){
//透视变换及裁剪
Image subImg = OcrUtils.transformAndCrop(srcMatList.get(i), boxList.get(i).get(j));
//高宽比 > 1.5 纵向
if (subImg.getHeight() * 1.0 / subImg.getWidth() > 1.5) {
//旋转图片90度
subImg = OcrUtils.rotateImg(manager, subImg);
isRotatedList.add(true);
imageList.add(subImg);
}else{
isRotatedList.add(false);
imageList.add(subImg);
}
index++;
}
}
List<List<OcrItem>> result = new ArrayList<>();
List<DirectionInfo> directionInfos = batchDetect(imageList);
if(CollectionUtils.isEmpty(directionInfos)){
throw new OcrException("方向检测失败");
}
index = 0;
for(int i = 0; i < srcMatList.size(); i++){
List<OcrItem> ocrItemList = new ArrayList<>();
for (int j = 0; j < boxList.get(i).size(); j++){
DirectionInfo directionInfo = directionInfos.get(index);
if(Objects.isNull(directionInfo)){
throw new OcrException("方向检测失败: 第" + i + "张图片, 第" + j + "个文本块,未检测到方向");
}
String angle;
if(isRotatedList.get(index)){
if (directionInfo.getName().equalsIgnoreCase("Rotate")) {
angle = "270";
} else {
angle = "90";
}
}else{
if (directionInfo.getName().equalsIgnoreCase("No Rotate")) {
angle = "0";
} else {
angle = "180";
}
}
OcrItem ocrItem = new OcrItem(boxList.get(i).get(j), AngleEnum.fromValue(angle), directionInfo.getProb().floatValue());
ocrItemList.add(ocrItem);
index++;
}
result.add(ocrItemList);
}
return result;
}
}
private List<DirectionInfo> batchDetect(List<Image> imageList) {
Predictor<Image, DirectionInfo> predictor = null;
try {
predictor = predictorPool.borrowObject();
return predictor.batchPredict(imageList);
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public void setTextDetModel(OcrCommonDetModel detModel) {
this.textDetModel = detModel;
}
@Override
public OcrCommonDetModel getTextDetModel() {
return textDetModel;
}
@Override
public void close() throws Exception {
try {
@@ -320,12 +385,5 @@ public class PPOCRMobileV2Model implements OcrDirectionModel {
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
try {
if (detModel != null) {
detModel.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,61 @@
package cn.smartjavaai.ocr.model.common.direction.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.ndarray.NDList;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.config.DirectionModelConfig;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.entity.DirectionInfo;
import cn.smartjavaai.ocr.enums.CommonDetModelEnum;
import cn.smartjavaai.ocr.enums.DirectionModelEnum;
import cn.smartjavaai.ocr.model.common.detect.translator.PPOCRDetTranslator;
import cn.smartjavaai.ocr.model.common.direction.translator.PpWordRotateTranslator;
import org.apache.commons.lang3.StringUtils;
import java.nio.file.Paths;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* 行方向分类
* @author dwj
*/
public class DirectionCriteriaFactory {
public static Criteria<Image, DirectionInfo> createCriteria(DirectionModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
}
Criteria<Image, DirectionInfo> criteria = null;
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
params.putAll(config.getCustomParams());
if(StringUtils.isNotBlank(config.getBatchifier())){
params.put("batchifier", config.getBatchifier());
}
if(config.getModelEnum() == DirectionModelEnum.CH_PPOCR_MOBILE_V2_CLS){
params.put("resizeWidth", 192);
params.put("resizeHeight", 48);
}else if (config.getModelEnum() == DirectionModelEnum.PP_LCNET_X0_25){
params.put("resizeWidth", 160);
params.put("resizeHeight", 80);
}else if (config.getModelEnum() == DirectionModelEnum.PP_LCNET_X1_0){
params.put("resizeWidth", 160);
params.put("resizeHeight", 80);
}
criteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, DirectionInfo.class)
.optModelPath(Paths.get(config.getModelPath()))
.optDevice(device)
.optTranslator(new PpWordRotateTranslator(params))
.optProgress(new ProgressBar())
.build();
return criteria;
}
}

View File

@@ -13,6 +13,7 @@ import cn.smartjavaai.ocr.entity.DirectionInfo;
import java.util.Arrays;
import java.util.List;
import java.util.Map;
/**
* 方向检测
@@ -24,7 +25,24 @@ import java.util.List;
public class PpWordRotateTranslator implements Translator<Image, DirectionInfo> {
List<String> classes = Arrays.asList("No Rotate", "Rotate");
public PpWordRotateTranslator() {
private String batchifier;
private int resizeHeight;
private int resizeWidth;
public PpWordRotateTranslator(Map<String, ?> arguments) {
batchifier = arguments.containsKey("batchifier")
? arguments.get("batchifier").toString()
: "padding";
resizeWidth = arguments.containsKey("resizeWidth")
? (Integer) arguments.get("resizeWidth")
: 192;
resizeHeight = arguments.containsKey("resizeHeight")
? (Integer) arguments.get("resizeHeight")
: 48;
}
@Override
@@ -51,8 +69,8 @@ public class PpWordRotateTranslator implements Translator<Image, DirectionInfo>
public NDList processInput(TranslatorContext ctx, Image input) {
NDArray img = input.toNDArray(ctx.getNDManager());
int imgC = 3;
int imgH = 48;
int imgW = 192;
int imgH = resizeHeight;
int imgW = resizeWidth;
NDArray array = ctx.getNDManager().zeros(new Shape(imgC, imgH, imgW));
@@ -74,13 +92,14 @@ public class PpWordRotateTranslator implements Translator<Image, DirectionInfo>
array.set(new NDIndex(":,:,0:" + resized_w), img);
array = array.expandDims(0);
// array = array.expandDims(0);
return new NDList(new NDArray[]{array});
}
@Override
public Batchifier getBatchifier() {
return null;
return Batchifier.fromString(batchifier);
}
}

View File

@@ -3,8 +3,11 @@ package cn.smartjavaai.ocr.model.common.recognize;
import ai.djl.modality.cv.Image;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.config.OcrRecModelConfig;
import cn.smartjavaai.ocr.config.OcrRecOptions;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.entity.OcrInfo;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.direction.OcrDirectionModel;
import java.awt.image.BufferedImage;
import java.util.List;
@@ -15,6 +18,22 @@ import java.util.List;
*/
public interface OcrCommonRecModel extends AutoCloseable{
default void setTextDetModel(OcrCommonDetModel detModel){
throw new UnsupportedOperationException("默认不支持该功能");
}
default OcrCommonDetModel getTextDetModel(){
throw new UnsupportedOperationException("默认不支持该功能");
}
default void setDirectionModel(OcrDirectionModel directionModel){
throw new UnsupportedOperationException("默认不支持该功能");
}
default OcrDirectionModel getDirectionModel(){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 加载模型
* @param config
@@ -26,7 +45,16 @@ public interface OcrCommonRecModel extends AutoCloseable{
* @param imagePath 图片路径
* @return
*/
default OcrInfo recognize(String imagePath) {
default OcrInfo recognize(String imagePath, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 文本识别
* @param image
* @return
*/
default OcrInfo recognize(Image image, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@@ -36,7 +64,7 @@ public interface OcrCommonRecModel extends AutoCloseable{
* @param image BufferedImage
* @return
*/
default OcrInfo recognize(BufferedImage image) {
default OcrInfo recognize(BufferedImage image, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@@ -46,7 +74,7 @@ public interface OcrCommonRecModel extends AutoCloseable{
* @param imageData 图片字节数组
* @return
*/
default OcrInfo recognize(byte[] imageData) {
default OcrInfo recognize(byte[] imageData, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@@ -56,7 +84,7 @@ public interface OcrCommonRecModel extends AutoCloseable{
* @param imagePath
* @param outputPath
*/
default void recognizeAndDraw(String imagePath, String outputPath, int fontSize) {
default void recognizeAndDraw(String imagePath, String outputPath, int fontSize, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@@ -65,7 +93,16 @@ public interface OcrCommonRecModel extends AutoCloseable{
* @param sourceImage
* @return
*/
default BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize){
default BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize, OcrRecOptions options){
throw new UnsupportedOperationException("默认不支持该功能");
}
default List<OcrInfo> batchRecognize(List<BufferedImage> imageList, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");
}
default List<OcrInfo> batchRecognizeDJLImage(List<Image> imageList, OcrRecOptions options) {
throw new UnsupportedOperationException("默认不支持该功能");
}

View File

@@ -6,35 +6,28 @@ import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.Point;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.opencv.OpenCVImageFactory;
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 ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.ocr.config.DirectionModelConfig;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.config.OcrRecModelConfig;
import cn.smartjavaai.ocr.config.OcrRecOptions;
import cn.smartjavaai.ocr.entity.*;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.factory.OcrModelFactory;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.detect.translator.PPOCRV5DetTranslator;
import cn.smartjavaai.ocr.model.common.direction.OcrDirectionModel;
import cn.smartjavaai.ocr.model.common.recognize.translator.PPOCRV5RecTranslator;
import cn.smartjavaai.ocr.opencv.OcrNDArrayUtils;
import cn.smartjavaai.ocr.model.common.recognize.criteria.OcrCommonRecCriterialFactory;
import cn.smartjavaai.ocr.opencv.OcrOpenCVUtils;
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;
@@ -45,50 +38,37 @@ import java.awt.image.BufferedImage;
import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.*;
import java.util.concurrent.ConcurrentHashMap;
import java.util.stream.Collectors;
/**
* PPOCRV5 识别模型
* @author dwj
* @date 2025/4/21
*/
@Slf4j
public class PpOCRV5RecModel implements OcrCommonRecModel {
public class OcrCommonRecModelImpl implements OcrCommonRecModel {
private ObjectPool<Predictor<Image, String>> recPredictorPool;
private OcrRecModelConfig config;
private OcrCommonDetModel detModel;
private ZooModel<Image, String> recognitionModel;
private OcrDirectionModel directionModel;
private ZooModel<Image, String> recognitionModel;
private OcrCommonDetModel textDetModel;
@Override
public void loadModel(OcrRecModelConfig config){
if(StringUtils.isBlank(config.getRecModelPath())){
throw new OcrException("recModelPath is null");
}
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
}
this.config = config;
this.directionModel = config.getDirectionModel();
this.textDetModel = config.getTextDetModel();
//初始化 识别Criteria
Criteria<Image, String> recCriteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, String.class)
.optModelPath(Paths.get(config.getRecModelPath()))
.optTranslator(new PPOCRV5RecTranslator(new ConcurrentHashMap<String, String>()))
.optProgress(new ProgressBar())
.optDevice(device)
.build();
Criteria<Image, String> recCriteria = OcrCommonRecCriterialFactory.createCriteria(config);
try{
recognitionModel = ModelZoo.loadModel(recCriteria);
this.recPredictorPool = new GenericObjectPool<>(new PredictorFactory<>(recognitionModel));
@@ -98,29 +78,11 @@ public class PpOCRV5RecModel implements OcrCommonRecModel {
throw new OcrException("识别模型加载失败", e);
}
//获取检测模型
if(StringUtils.isNotBlank(config.getDetModelPath()) && Objects.nonNull(config.getDetModelEnum())){
OcrDetModelConfig detModelConfig = new OcrDetModelConfig();
detModelConfig.setModelEnum(config.getDetModelEnum());
detModelConfig.setDetModelPath(config.getDetModelPath());
detModelConfig.setDevice(config.getDevice());
detModel = OcrModelFactory.getInstance().getDetModel(detModelConfig);
}
//获取方向检测模型
if(StringUtils.isNotBlank(config.getDirectionModelPath()) && Objects.nonNull(config.getDirectionModelEnum())){
DirectionModelConfig directionModelConfig = new DirectionModelConfig();
directionModelConfig.setModelEnum(config.getDirectionModelEnum());
directionModelConfig.setModelPath(config.getDirectionModelPath());
directionModelConfig.setDevice(config.getDevice());
directionModel = OcrModelFactory.getInstance().getDirectionModel(directionModelConfig);
}
}
@Override
public OcrInfo recognize(String imagePath) {
public OcrInfo recognize(String imagePath, OcrRecOptions options) {
if(StringUtils.isBlank(config.getRecModelPath())){
throw new OcrException("recModelPath为空无法识别");
}
@@ -130,78 +92,44 @@ public class PpOCRV5RecModel implements OcrCommonRecModel {
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
return recognize(img, options);
} catch (IOException e) {
throw new OcrException("无效的图片", e);
}
OcrInfo ocrInfo = recognize(img);
((Mat)img.getWrappedImage()).release();
return ocrInfo;
}
private OcrInfo recognize(Image image) {
//检测文本
List<OcrBox> boxeList = detModel.detect(image);
if(Objects.isNull(boxeList) || boxeList.isEmpty()){
throw new OcrException("未检测到文本");
}
Predictor<Image, String> predictor = null;
List<RotatedBox> rotatedBoxes = new ArrayList<>();
List<OcrItem> ocrItemList = new ArrayList<>();
try (NDManager manager = NDManager.newBaseManager()) {
Mat srcMat = (Mat) image.getWrappedImage();
predictor = recPredictorPool.borrowObject();
//检测方向
if(directionModel != null){
ocrItemList = directionModel.detect(boxeList, srcMat, manager);
if(Objects.isNull(ocrItemList) || ocrItemList.isEmpty()){
throw new OcrException("方向检测失败");
}
for (OcrItem ocrItem : ocrItemList){
//放射变换+裁剪
Image subImage = OcrUtils.transformAndCrop(srcMat, ocrItem.getOcrBox());
//ImageUtils.saveImage(subImage, UUID.randomUUID().toString() + "_aaa.png", "build/output");
//纠正文本框
subImage = OcrUtils.rotateImg(subImage, ocrItem.getAngle());
//ImageUtils.saveImage(subImage, UUID.randomUUID().toString() + "_bbb.png", "build/output");
//识别
String name = predictor.predict(subImage);
ocrItem.setText(name);
NDArray ndArray = manager.create(ocrItem.getOcrBox().toFloatArray());
rotatedBoxes.add(new RotatedBox(ndArray, ocrItem.getText()));
((Mat)subImage.getWrappedImage()).release();
}
}else{
for (OcrBox box : boxeList){
RotatedBox rotatedBox = recognize(box, srcMat, predictor, manager);
rotatedBoxes.add(rotatedBox);
}
}
//后处理
return postProcessOcrResult(rotatedBoxes);
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
recPredictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
} finally {
if(img != null){
((Mat)img.getWrappedImage()).release();
}
}
}
/**
*
* @param image
* @param options
* @return
*/
@Override
public OcrInfo recognize(Image image, OcrRecOptions options) {
List<OcrInfo> result = batchRecognizeDJLImage(Collections.singletonList(image), options);
if(CollectionUtils.isEmpty(result)){
throw new OcrException("OCR识别结果为空");
}
return result.get(0);
}
private RotatedBox recognize(OcrBox box,Mat srcMat,Predictor<Image, String> recPredictor,NDManager manager){
try {
/**
* 批量矫正文本框
* @param boxList
* @param srcMat
* @param manager
* @return
*/
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++) {
//透视变换 + 裁剪
Image subImg = OcrUtils.transformAndCrop(srcMat, box);
Image subImg = OcrUtils.transformAndCrop(srcMat, boxList.get(i));
//ImageUtils.saveImage(subImg, i + "crop.png", "build/output");
//高宽比 > 1.5
if (subImg.getHeight() * 1.0 / subImg.getWidth() > 1.5) {
@@ -209,21 +137,63 @@ public class PpOCRV5RecModel implements OcrCommonRecModel {
subImg = OcrUtils.rotateImg(manager, subImg);
//ImageUtils.saveImage(subImg, i + "rotate.png", "build/output");
}
String name = recPredictor.predict(subImg);
((Mat)subImg.getWrappedImage()).release();
NDArray pointsArray = manager.create(box.toFloatArray());
return new RotatedBox(pointsArray, name);
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
imageList.add(subImg);
}
return imageList;
}
/**
* 批量矫正文本框
* @param itemList
* @param srcMat
* @param manager
* @return
*/
private List<Image> batchAlignWithDirection(List<OcrItem> itemList, Mat srcMat,NDManager manager){
List<Image> imageList = new ArrayList<>(itemList.size());
for (OcrItem ocrItem : itemList) {
//放射变换+裁剪
Image subImage = OcrUtils.transformAndCrop(srcMat, ocrItem.getOcrBox());
//ImageUtils.saveImage(subImage, UUID.randomUUID().toString() + "_aaa.png", "build/output");
//纠正文本框
subImage = OcrUtils.rotateImg(subImage, ocrItem.getAngle());
imageList.add(subImage);
}
return imageList;
}
// private RotatedBox recognize(OcrBox box,Mat srcMat,Predictor<Image, String> recPredictor,NDManager manager){
// try {
// //透视变换 + 裁剪
// Image subImg = OcrUtils.transformAndCrop(srcMat, box);
// //ImageUtils.saveImage(subImg, i + "crop.png", "build/output");
// //高宽比 > 1.5
// if (subImg.getHeight() * 1.0 / subImg.getWidth() > 1.5) {
// //旋转图片90度
// subImg = OcrUtils.rotateImg(manager, subImg);
// //ImageUtils.saveImage(subImg, i + "rotate.png", "build/output");
// }
// String name = recPredictor.predict(subImg);
// ((Mat)subImg.getWrappedImage()).release();
// NDArray pointsArray = manager.create(box.toFloatArray());
// return new RotatedBox(pointsArray, name);
// } catch (Exception e) {
// throw new OcrException("OCR检测错误", e);
// }
// }
/**
* 后处理排序,分行
* @param rotatedBoxes
*/
private OcrInfo postProcessOcrResult(List<RotatedBox> rotatedBoxes){
private OcrInfo postProcessOcrResult(List<RotatedBox> rotatedBoxes, OcrRecOptions ocrRecOptions){
//不分行
if(!ocrRecOptions.isEnableLineSplit()){
return OcrUtils.convertRotatedBoxesToOcrItems(rotatedBoxes);
}
//Y坐标升序排序
List<RotatedBox> initList = new ArrayList<>();
for (RotatedBox result : rotatedBoxes) {
@@ -257,13 +227,13 @@ public class PpOCRV5RecModel implements OcrCommonRecModel {
@Override
public void recognizeAndDraw(String imagePath, String outputPath, int fontSize) {
public void recognizeAndDraw(String imagePath, String outputPath, int fontSize, OcrRecOptions options) {
if(!FileUtils.isFileExists(imagePath)){
throw new OcrException("图像文件不存在");
}
try {
Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
OcrInfo ocrInfo = recognize(img);
OcrInfo ocrInfo = recognize(img, options);
if(Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()){
throw new OcrException("未检测到文字");
}
@@ -278,36 +248,36 @@ public class PpOCRV5RecModel implements OcrCommonRecModel {
}
@Override
public OcrInfo recognize(BufferedImage image) {
public OcrInfo recognize(BufferedImage image, OcrRecOptions options) {
if(!ImageUtils.isImageValid(image)){
throw new OcrException("图像无效");
}
Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
OcrInfo ocrInfo = recognize(img);
OcrInfo ocrInfo = recognize(img, options);
((Mat)img.getWrappedImage()).release();
return ocrInfo;
}
@Override
public OcrInfo recognize(byte[] imageData) {
public OcrInfo recognize(byte[] imageData, OcrRecOptions options) {
if(Objects.isNull(imageData)){
throw new OcrException("图像无效");
}
try {
BufferedImage image = ImageIO.read(new ByteArrayInputStream(imageData));
return recognize(image);
return recognize(image, options);
} catch (IOException e) {
throw new OcrException("错误的图像", e);
}
}
@Override
public BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize) {
public BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize, OcrRecOptions options) {
if(!ImageUtils.isImageValid(sourceImage)){
throw new OcrException("图像无效");
}
Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage));
OcrInfo ocrInfo = recognize(img);
OcrInfo ocrInfo = recognize(img, options);
if(Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()){
throw new OcrException("未检测到文字");
}
@@ -325,6 +295,154 @@ public class PpOCRV5RecModel implements OcrCommonRecModel {
}
}
@Override
public List<OcrInfo> batchRecognize(List<BufferedImage> imageList, OcrRecOptions options) {
List<Image> djlImageList = new ArrayList<>(imageList.size());
try {
for (BufferedImage bufferedImage : imageList) {
djlImageList.add(ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(bufferedImage)));
}
return batchRecognizeDJLImage(djlImageList, options);
} catch (Exception e) {
throw new OcrException(e);
} finally {
djlImageList.forEach(image -> ((Mat)image.getWrappedImage()).release());
}
}
@Override
public List<OcrInfo> batchRecognizeDJLImage(List<Image> imageList, OcrRecOptions options) {
if(Objects.isNull(textDetModel)){
throw new OcrException("textDetModel is null");
}
OcrRecOptions ocrRecOptions = options;
if(Objects.isNull(options)){
ocrRecOptions = new OcrRecOptions();
}
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()){
throw new OcrException("未检测到文本");
}
Predictor<Image, String> predictor = null;
List<OcrInfo> ocrInfoList = new ArrayList<OcrInfo>();
try (NDManager manager = NDManager.newBaseManager()) {
predictor = recPredictorPool.borrowObject();
List<Image> allImageAlignList = new ArrayList<Image>();
//检测方向
if(ocrRecOptions.isEnableDirectionCorrect()){
if(Objects.isNull(directionModel)){
throw new OcrException("请配置方向模型");
}
List<Mat> matList = imageList.stream()
.map(image -> (Mat)image.getWrappedImage())
.collect(Collectors.toList());
List<List<OcrItem>> ocrItemList = directionModel.batchDetect(boxeList, matList);
if(CollectionUtils.isEmpty(ocrItemList) || ocrItemList.size() != imageList.size()){
throw new OcrException("方向检测失败");
}
allImageAlignList = new ArrayList<Image>();
for (int i = 0; i < ocrItemList.size(); i++) {
Mat srcMat = (Mat) imageList.get(i).getWrappedImage();
List<Image> imageAlignList = batchAlignWithDirection(ocrItemList.get(i), srcMat, manager);
// for(int j = 0; j < imageAlignList.size(); j++){
// ImageUtils.saveImage(imageAlignList.get(j),"dir-"+i+"-"+j+".png","/Users/xxx/Downloads/testing33");
// }
allImageAlignList.addAll(imageAlignList);
}
}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);
// for(int j = 0; j < imageAlignList.size(); j++){
// ImageUtils.saveImage(imageAlignList.get(j),i+"-"+j+".png","/Users/xxx/Downloads/testing33");
// }
allImageAlignList.addAll(imageAlignList);
}
}
List<String> textList = batchRecognize(allImageAlignList);
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()){
throw new OcrException("识别失败: 第" + i + "张图片, 第" + j + "个文本块,未识别到文本");
}
OcrBox box = boxeList.get(i).get(j);
NDArray pointsArray = manager.create(box.toFloatArray());
rotatedBoxes.add(new RotatedBox(pointsArray, textList.get(textIndex)));
textIndex++;
}
OcrInfo ocrInfo = postProcessOcrResult(rotatedBoxes, ocrRecOptions);
ocrInfoList.add(ocrInfo);
}
return ocrInfoList;
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
recPredictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
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());
return textList;
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
recPredictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public void setTextDetModel(OcrCommonDetModel detModel) {
this.textDetModel = detModel;
}
@Override
public OcrCommonDetModel getTextDetModel() {
return textDetModel;
}
@Override
public void setDirectionModel(OcrDirectionModel directionModel) {
this.directionModel = directionModel;
}
@Override
public OcrDirectionModel getDirectionModel() {
return directionModel;
}
@Override
public void close() throws Exception {
try {
@@ -334,20 +452,6 @@ public class PpOCRV5RecModel implements OcrCommonRecModel {
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (detModel != null) {
detModel.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
try {
if (directionModel != null) {
directionModel.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
try {
if (recognitionModel != null) {
recognitionModel.close();

View File

@@ -0,0 +1,51 @@
package cn.smartjavaai.ocr.model.common.recognize.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.config.OcrRecModelConfig;
import cn.smartjavaai.ocr.enums.CommonRecModelEnum;
import cn.smartjavaai.ocr.model.common.recognize.translator.PPOCRRecTranslator;
import org.apache.commons.lang3.StringUtils;
import java.nio.file.Paths;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* @author dwj
* @date 2025/7/8
*/
public class OcrCommonRecCriterialFactory {
public static Criteria<Image, String> createCriteria(OcrRecModelConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
}
Criteria<Image, String> criteria = null;
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
params.putAll(config.getCustomParams());
if(StringUtils.isNotBlank(config.getBatchifier())){
params.put("batchifier", config.getBatchifier());
}
if(config.getRecModelEnum() == CommonRecModelEnum.PP_OCR_V5_SERVER_REC_MODEL ||
config.getRecModelEnum() == CommonRecModelEnum.PP_OCR_V5_MOBILE_REC_MODEL ||
config.getRecModelEnum() == CommonRecModelEnum.PP_OCR_V4_SERVER_REC_MODEL ||
config.getRecModelEnum() == CommonRecModelEnum.PP_OCR_V4_MOBILE_REC_MODEL ){
criteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, String.class)
.optModelPath(Paths.get(config.getRecModelPath()))
.optTranslator(new PPOCRRecTranslator(params))
.optProgress(new ProgressBar())
.optDevice(device)
.build();
}
return criteria;
}
}

View File

@@ -23,15 +23,20 @@ import java.util.Map;
* 文字识别前后处理
*
*/
public class PPOCRV5RecTranslator implements Translator<Image, String> {
public class PPOCRRecTranslator implements Translator<Image, String> {
private List<String> table;
private final boolean use_space_char;
public PPOCRV5RecTranslator(Map<String, ?> arguments) {
private String batchifier;
public PPOCRRecTranslator(Map<String, ?> arguments) {
use_space_char =
arguments.containsKey("use_space_char")
? Boolean.parseBoolean(arguments.get("use_space_char").toString())
: true;
batchifier = arguments.containsKey("batchifier")
? arguments.get("batchifier").toString()
: "padding";
}
@Override
@@ -57,7 +62,8 @@ public class PPOCRV5RecTranslator implements Translator<Image, String> {
StringBuilder sb = new StringBuilder();
NDArray tokens = list.singletonOrThrow();
long[] indices = tokens.get(0).argMax(1).toLongArray();
// long[] indices = tokens.get(0).argMax(1).toLongArray();
long[] indices = tokens.argMax(1).toLongArray();
boolean[] selection = new boolean[indices.length];
Arrays.fill(selection, true);
for (int i = 1; i < indices.length; i++) {
@@ -111,13 +117,13 @@ public class PPOCRV5RecTranslator implements Translator<Image, String> {
padding_im.set(new NDIndex(":,:,0:" + resized_w), resized_image);
padding_im = padding_im.flip(0);
padding_im = padding_im.expandDims(0);
// padding_im = padding_im.expandDims(0);
return new NDList(padding_im);
}
@Override
public Batchifier getBatchifier() {
return null;
return Batchifier.fromString(batchifier);
}
}

View File

@@ -0,0 +1,160 @@
package cn.smartjavaai.ocr.model.table;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
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.entity.R;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.config.TableStructureConfig;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.entity.OcrItem;
import cn.smartjavaai.ocr.entity.TableStructureResult;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.model.table.criteria.StructureCriteriaFactory;
import cn.smartjavaai.ocr.utils.OcrUtils;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.ObjectPool;
import org.apache.commons.pool2.impl.GenericObjectPool;
import org.opencv.core.Mat;
import javax.imageio.ImageIO;
import java.awt.image.BufferedImage;
import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.nio.file.Paths;
import java.util.List;
import java.util.Objects;
/**
* 表格结构模型
* @author dwj
*/
@Slf4j
public class CommonTableStructureModel implements TableStructureModel{
private ZooModel<Image, TableStructureResult> model;
private ObjectPool<Predictor<Image, TableStructureResult>> predictorPool;
@Override
public void loadModel(TableStructureConfig config) {
if(StringUtils.isBlank(config.getModelPath())){
throw new OcrException("modelPath is null");
}
Criteria<Image, TableStructureResult> criteria = StructureCriteriaFactory.createCriteria(config);
try{
model = ModelZoo.loadModel(criteria);
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model));
log.debug("当前设备: " + model.getNDManager().getDevice());
log.debug("当前引擎: " + Engine.getInstance().getEngineName());
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new OcrException("表格结构识别模型加载失败", e);
}
}
@Override
public R<TableStructureResult> detect(BufferedImage image) {
if(!ImageUtils.isImageValid(image)){
return R.fail(R.Status.INVALID_IMAGE);
}
Image img = null;
try {
img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
return detect(img);
} catch (Exception e) {
throw new OcrException(e);
} finally {
if(Objects.nonNull(img)){
((Mat)img.getWrappedImage()).release();
}
}
}
@Override
public R<TableStructureResult> detect(String imagePath) {
if(!FileUtils.isFileExists(imagePath)){
return R.fail(R.Status.FILE_NOT_FOUND);
}
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
return detect(img);
} catch (IOException e) {
throw new OcrException("无效的图片", e);
} finally {
if (Objects.nonNull(img)){
((Mat)img.getWrappedImage()).release();
}
}
}
@Override
public R<TableStructureResult> detect(byte[] imageData) {
if(Objects.isNull(imageData)){
return R.fail(R.Status.INVALID_IMAGE);
}
try {
BufferedImage image = ImageIO.read(new ByteArrayInputStream(imageData));
return detect(image);
} catch (IOException e) {
throw new OcrException("错误的图像", e);
}
}
@Override
public R<TableStructureResult> detect(Image image) {
Predictor<Image, TableStructureResult> predictor = null;
try {
predictor = predictorPool.borrowObject();
TableStructureResult result = predictor.predict(image);
return R.ok(result);
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public void close() throws Exception {
try {
if (predictorPool != null) {
predictorPool.close();
}
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (model != null) {
model.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}

View File

@@ -0,0 +1,454 @@
package cn.smartjavaai.ocr.model.table;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.output.Rectangle;
import ai.djl.translate.TranslateException;
import ai.djl.util.JsonUtils;
import cn.smartjavaai.common.entity.DetectionRectangle;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.ocr.config.OcrRecModelConfig;
import cn.smartjavaai.ocr.config.OcrRecOptions;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.entity.OcrInfo;
import cn.smartjavaai.ocr.entity.OcrItem;
import cn.smartjavaai.ocr.entity.TableStructureResult;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel;
import cn.smartjavaai.ocr.model.common.direction.OcrDirectionModel;
import cn.smartjavaai.ocr.model.common.recognize.OcrCommonRecModel;
import cn.smartjavaai.ocr.utils.ConvertHtml2Excel;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.lang3.tuple.Pair;
import org.apache.poi.hssf.usermodel.HSSFWorkbook;
import org.opencv.core.Mat;
import javax.imageio.ImageIO;
import java.awt.*;
import java.awt.image.BufferedImage;
import java.io.ByteArrayInputStream;
import java.io.File;
import java.io.IOException;
import java.nio.file.Paths;
import java.util.*;
import java.util.List;
import java.util.concurrent.ConcurrentHashMap;
/**
* 表格内容识别器
* @author dwj
*/
@Slf4j
public class TableRecognizer {
private OcrCommonDetModel textDetector;
private TableStructureModel tableStructureModel;
private OcrCommonRecModel textRecModel;
private OcrDirectionModel directionModel;
private TableRecognizer(Builder builder) {
this.tableStructureModel = builder.tableStructureModel;
this.textRecModel = builder.textRecModel;
this.directionModel = builder.directionModel;
this.textDetector = builder.textDetector;
textRecModel.setTextDetModel(textDetector);
textRecModel.setDirectionModel(directionModel);
}
public static Builder builder() {
return new Builder();
}
// 链式设置文本识别模型
public TableRecognizer withTextRecModel(OcrCommonRecModel textRecModel) {
this.textRecModel = textRecModel;
return this;
}
// 链式设置表格结构模型
public TableRecognizer withStructureModel(TableStructureModel tableStructureModel) {
this.tableStructureModel = tableStructureModel;
return this;
}
/**
* 表格识别
* @param image
* @return
*/
public R<TableStructureResult> recognize(Image image) {
//表格结构识别
R<TableStructureResult> result = tableStructureModel.detect(image);
if(!result.isSuccess()){
return R.fail(result.getCode(), result.getMessage());
}
//文本检测+文字识别
boolean enableDirectionCorrect = directionModel == null ? false : true;
OcrRecOptions options = new OcrRecOptions(enableDirectionCorrect, false);
OcrInfo ocrInfo = textRecModel.recognize(image, options);
List<String> tableContentList = buildTable(result.getData(), ocrInfo);
String html = convertHtml(result.getData().getTableTagList(), tableContentList);
result.getData().setHtml(html);
return result;
}
/**
* 表格识别
* @param image
* @return
*/
public R<TableStructureResult> recognize(BufferedImage image) {
if(!ImageUtils.isImageValid(image)){
return R.fail(R.Status.INVALID_IMAGE);
}
Image img = null;
try {
img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
return recognize(img);
} catch (Exception e) {
throw new OcrException(e);
} finally {
if(Objects.nonNull(img)){
((Mat)img.getWrappedImage()).release();
}
}
}
/**
* 表格识别
* @param imagePath
* @return
*/
public R<TableStructureResult> recognize(String imagePath) {
if(!FileUtils.isFileExists(imagePath)){
return R.fail(R.Status.FILE_NOT_FOUND);
}
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
return recognize(img);
} catch (IOException e) {
throw new OcrException("无效的图片", e);
} finally {
if (Objects.nonNull(img)){
((Mat)img.getWrappedImage()).release();
}
}
}
/**
* 表格识别
* @param imageData
* @return
*/
public R<TableStructureResult> recognize(byte[] imageData) {
if(Objects.isNull(imageData)){
return R.fail(R.Status.INVALID_IMAGE);
}
try {
BufferedImage image = ImageIO.read(new ByteArrayInputStream(imageData));
return recognize(image);
} catch (IOException e) {
throw new OcrException("错误的图像", e);
}
}
/**
* 绘制表格
* @param tableStructureResult
* @param image
* @param savePath
*/
public void drawTable(TableStructureResult tableStructureResult, BufferedImage image, String savePath){
if(Objects.isNull(tableStructureResult) || CollectionUtils.isEmpty(tableStructureResult.getTableTagList())){
throw new OcrException("表格结构为空");
}
for (int i = 0; i < tableStructureResult.getOcrItemList().size(); i++){
OcrItem item = tableStructureResult.getOcrItemList().get(i);
DetectionRectangle detectionRectangle = item.getOcrBox().toDetectionRectangle();
ImageUtils.drawImageRectWithText(image, detectionRectangle, i + "", Color.RED);
}
ImageUtils.saveImage(image, savePath);
}
/**
* 删除 HTML 中第一个 <style> ... </style> 段落
* @param html 原始 HTML
* @return 去掉 <style> 的 HTML
*/
public static String removeStyleBlock(String html) {
String lowerHtml = html.toLowerCase();
int styleStart = lowerHtml.indexOf("<style");
if (styleStart == -1) {
return html; // 没有 style返回原文
}
int styleEnd = lowerHtml.indexOf("</style>", styleStart);
if (styleEnd == -1) {
return html; // 没闭合标签,不处理
}
styleEnd += "</style>".length();
// 去掉 style 块
return html.substring(0, styleStart) + html.substring(styleEnd);
}
/**
* 导出 Excel
* @param html
* @param savePath
*/
public void exportExcel(String html, String savePath){
try {
String content = removeStyleBlock(html);
content = content.replace("<html><body>", "");
content = content.replace("</body></html>", "");
HSSFWorkbook workbook = ConvertHtml2Excel.table2Excel(content);
workbook.write(new File(savePath));
} catch (Exception e) {
throw new OcrException("导出excel失败请检查表结构是否识别正确");
}
}
/**
* 构建表格
* @param tableStructureResult
* @param ocrInfo
* @return
*/
public List<String> buildTable(TableStructureResult tableStructureResult, OcrInfo ocrInfo) {
// 获取 Cell 与 文本检测框 的对应关系(1:N)。
Map<Integer, List<Integer>> matched = new ConcurrentHashMap<>();
List<OcrItem> ocrItems = ocrInfo.getOcrItemList();
for (int i = 0; i < ocrItems.size(); i++) {
OcrBox ocrBox = ocrItems.get(i).getOcrBox();
int[] box_1 = {
(int)ocrBox.getTopLeft().getX(),
(int)ocrBox.getTopLeft().getY(),
(int)ocrBox.getBottomRight().getX(),
(int)ocrBox.getBottomRight().getY()
};
// 获取两两cell之间的L1距离和 1- IOU
List<Pair<Float, Float>> distances = new ArrayList<>();
for (OcrItem cell : tableStructureResult.getOcrItemList()) {
OcrBox cellBox = cell.getOcrBox();
int[] box_2 = {
(int)cellBox.getTopLeft().getX(),
(int)cellBox.getTopLeft().getY(),
(int)cellBox.getBottomRight().getX(),
(int)cellBox.getBottomRight().getY()
};
float distance = distance(box_1, box_2);
float iou = 1 - computeIou(box_1, box_2);
distances.add(Pair.of(distance, iou));
}
// 根据距离和IOU挑选最"近"的cell
Pair<Float, Float> nearest = sorted(distances);
// 获取最小距离对应的下标id也等价于cell的下标id distances列表是根据遍历cells生成的
int id = 0;
for (int idx = 0; idx < distances.size(); idx++) {
Pair<Float, Float> current = distances.get(idx);
if (current.getLeft().floatValue() == nearest.getLeft().floatValue()
&& current.getRight().floatValue() == nearest.getRight().floatValue()) {
id = idx;
break;
}
}
if (!matched.containsKey(id)) {
List<Integer> textIds = new ArrayList<>();
textIds.add(i);
// cell id, text id list (dt_boxes index list)
matched.put(id, textIds);
} else {
matched.get(id).add(i);
}
}
List<String> cell_contents = new ArrayList<>();
List<Double> probs = new ArrayList<>();
for (int i = 0; i < tableStructureResult.getOcrItemList().size(); i++) {
List<Integer> textIds = matched.get(i);
List<String> contents = new ArrayList<>();
String content = "";
if (textIds != null) {
for (Integer id : textIds) {
contents.add(ocrItems.get(id).getText());
}
content = StringUtils.join(contents, " ");
}
cell_contents.add(content);
probs.add(-1.0);
}
return cell_contents;
}
/**
* 计算欧式距离
* Calculate L1 distance
*
* @param box_1
* @param box_2
* @return
*/
private int distance(int[] box_1, int[] box_2) {
int x1 = box_1[0];
int y1 = box_1[1];
int x2 = box_1[2];
int y2 = box_1[3];
int x3 = box_2[0];
int y3 = box_2[1];
int x4 = box_2[2];
int y4 = box_2[3];
int dis = Math.abs(x3 - x1) + Math.abs(y3 - y1) + Math.abs(x4 - x2) + Math.abs(y4 - y2);
int dis_2 = Math.abs(x3 - x1) + Math.abs(y3 - y1);
int dis_3 = Math.abs(x4 - x2) + Math.abs(y4 - y2);
return dis + Math.min(dis_2, dis_3);
}
/**
* 计算交并比
* computing IoU
*
* @param rec1: (y0, x0, y1, x1), which reflects (top, left, bottom, right)
* @param rec2: (y0, x0, y1, x1)
* @return scala value of IoU
*/
private float computeIou(int[] rec1, int[] rec2) {
// computing area of each rectangles
int S_rec1 = (rec1[2] - rec1[0]) * (rec1[3] - rec1[1]);
int S_rec2 = (rec2[2] - rec2[0]) * (rec2[3] - rec2[1]);
// computing the sum_area
int sum_area = S_rec1 + S_rec2;
// find the each edge of intersect rectangle
int left_line = Math.max(rec1[1], rec2[1]);
int right_line = Math.min(rec1[3], rec2[3]);
int top_line = Math.max(rec1[0], rec2[0]);
int bottom_line = Math.min(rec1[2], rec2[2]);
// judge if there is an intersect
if (left_line >= right_line || top_line >= bottom_line) {
return 0.0f;
} else {
float intersect = (right_line - left_line) * (bottom_line - top_line);
return (intersect / (sum_area - intersect)) * 1.0f;
}
}
/**
* 距离排序
* Distance sorted
*
* @param distances
* @return
*/
private Pair<Float, Float> sorted(List<Pair<Float, Float>> distances) {
Comparator<Pair<Float, Float>> comparator =
new Comparator<Pair<Float, Float>>() {
@Override
public int compare(Pair<Float, Float> a1, Pair<Float, Float> a2) {
// 首先根据IoU排序
if (a1.getRight().floatValue() > a2.getRight().floatValue()) {
return 1;
} else if (a1.getRight().floatValue() == a2.getRight().floatValue()) {
// 然后根据L1距离排序
if (a1.getLeft().floatValue() > a2.getLeft().floatValue()) {
return 1;
}
return -1;
}
return -1;
}
};
// 距离排序
List<Pair<Float, Float>> newDistances = new ArrayList<>();
CollectionUtils.addAll(newDistances, new Object[distances.size()]);
Collections.copy(newDistances, distances);
Collections.sort(newDistances, comparator);
return newDistances.get(0);
}
/**
* 生成表格html
* Generate table html
*
* @param pred_structures
* @param cell_contents
* @return
*/
public String convertHtml(List<String> pred_structures, List<String> cell_contents) {
StringBuffer html = new StringBuffer();
// 添加统一的样式(可选放到<head>中)
html.append("<style>\n");
html.append("table { border-collapse: collapse; }\n");
html.append("td, th, table { border: 1px solid black; padding: 5px; }\n");
html.append("</style>\n");
int td_index = 0;
for (String tag : pred_structures) {
if (tag.contains("<td></td>")) {
String content = cell_contents.get(td_index);
html.append("<td>");
html.append(content);
html.append("</td>");
td_index++;
continue;
}
html.append(tag);
}
return html.toString();
}
public static class Builder {
private TableStructureModel tableStructureModel;
private OcrCommonRecModel textRecModel;
private OcrDirectionModel directionModel;
private OcrCommonDetModel textDetector;
public Builder withStructureModel(TableStructureModel model) {
this.tableStructureModel = model;
return this;
}
public Builder withTextRecModel(OcrCommonRecModel model) {
this.textRecModel = model;
return this;
}
public Builder withDirectionModel(OcrDirectionModel model) {
this.directionModel = model;
return this;
}
public Builder withTextDetModel(OcrCommonDetModel model) {
this.textDetector = model;
return this;
}
public TableRecognizer build() {
if (this.tableStructureModel == null) {
throw new IllegalStateException("tableStructureModel 未设置");
}
if (this.textDetector == null) {
throw new IllegalStateException("textDetector 未设置");
}
if (this.textRecModel == null) {
throw new IllegalStateException("textRecModel 未设置");
}
return new TableRecognizer(this);
}
}
}

View File

@@ -0,0 +1,63 @@
package cn.smartjavaai.ocr.model.table;
import ai.djl.modality.cv.Image;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.ocr.config.OcrDetModelConfig;
import cn.smartjavaai.ocr.config.TableStructureConfig;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.entity.OcrItem;
import cn.smartjavaai.ocr.entity.TableStructureResult;
import java.awt.image.BufferedImage;
import java.util.List;
/**
* 表格结构识别模型
* @author dwj
*/
public interface TableStructureModel extends AutoCloseable{
/**
* 加载模型
* @param config
*/
void loadModel(TableStructureConfig config);
/**
* 表格结构检测
* @param image
* @return
*/
default R<TableStructureResult> detect(BufferedImage image){
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 表格结构检测
* @param imagePath 图片路径
* @return
*/
default R<TableStructureResult> detect(String imagePath) {
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 表格结构检测
* @param imageData 图片字节数组
* @return
*/
default R<TableStructureResult> detect(byte[] imageData) {
throw new UnsupportedOperationException("默认不支持该功能");
}
/**
* 表格结构检测
* @param image DJL Image
* @return
*/
default R<TableStructureResult> detect(Image image){
throw new UnsupportedOperationException("默认不支持该功能");
}
}

View File

@@ -0,0 +1,60 @@
package cn.smartjavaai.ocr.model.table.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.ocr.config.TableStructureConfig;
import cn.smartjavaai.ocr.entity.OcrItem;
import cn.smartjavaai.ocr.entity.TableStructureResult;
import cn.smartjavaai.ocr.enums.TableStructureModelEnum;
import cn.smartjavaai.ocr.model.table.translator.TableStructTranslator;
import java.nio.file.Paths;
import java.util.List;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* @author dwj
* @date 2025/7/10
*/
public class StructureCriteriaFactory {
public static Criteria<Image, TableStructureResult> createCriteria(TableStructureConfig config) {
Device device = null;
if(!Objects.isNull(config.getDevice())){
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Criteria<Image, TableStructureResult> criteria = null;
if(config.getModelEnum() == TableStructureModelEnum.SLANET){
criteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, TableStructureResult.class)
.optModelPath(Paths.get(config.getModelPath()))
.optOption("removePass", "repeated_fc_relu_fuse_pass")
.optDevice(device)
.optTranslator(new TableStructTranslator())
.optProgress(new ProgressBar())
.build();
}else if(config.getModelEnum() == TableStructureModelEnum.SLANET_PLUS){
criteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, TableStructureResult.class)
.optModelPath(Paths.get(config.getModelPath()))
.optOption("removePass", "repeated_fc_relu_fuse_pass")
.optDevice(device)
.optTranslator(new TableStructTranslator())
.optProgress(new ProgressBar())
.build();
}
return criteria;
}
}

View File

@@ -0,0 +1,198 @@
package cn.smartjavaai.ocr.model.table.translator;
import ai.djl.Model;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.Rectangle;
import ai.djl.modality.cv.util.NDImageUtils;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.index.NDIndex;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import ai.djl.translate.Batchifier;
import ai.djl.translate.Translator;
import ai.djl.translate.TranslatorContext;
import ai.djl.util.Utils;
import cn.smartjavaai.common.entity.Point;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.entity.OcrItem;
import cn.smartjavaai.ocr.entity.TableStructureResult;
import java.io.IOException;
import java.io.InputStream;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
/**
* 表格识别的前后处理
*/
public class TableStructTranslator implements Translator<Image, TableStructureResult> {
private final int maxLength = 488;
private int height;
private int width;
private float scale = 1.0f;
private float xScale;
private float yScale;
private List<String> dict;
private String beg_str = "sos";
private String end_str = "eos";
private List<String> td_token = new ArrayList<>();
@Override
public void prepare(TranslatorContext ctx) throws IOException {
Model model = ctx.getModel();
try (InputStream is = model.getArtifact("table_structure_dict_ch.txt").openStream()) {
dict = Utils.readLines(is, false);
dict.add(0,beg_str);
if(dict.contains("<td>"))
dict.remove("<td>");
if(!dict.contains("<td></td>"))
dict.add("<td></td>");
dict.add(end_str);
}
td_token.add("<td>");
td_token.add("<td");
td_token.add("<td></td>");
}
@Override
public NDList processInput(TranslatorContext ctx, Image input) {
NDArray img = input.toNDArray(ctx.getNDManager(), Image.Flag.COLOR);
height = input.getHeight();
width = input.getWidth();
img = ResizeTableImage(img, height, width, maxLength);
img = PaddingTableImage(ctx, img, maxLength);
img = img.transpose(2, 0, 1).div(255).flip(0);
img = NDImageUtils.normalize(
img, new float[]{0.485f, 0.456f, 0.406f}, new float[]{0.229f, 0.224f, 0.225f});
img = img.expandDims(0);
return new NDList(img);
}
@Override
public TableStructureResult processOutput(TranslatorContext ctx, NDList list) {
NDArray bbox_preds = list.get(0);
NDArray structure_probs = list.get(1);
NDArray structure_idx = structure_probs.argMax(2);
structure_probs = structure_probs.max(new int[]{2});
List<List<String>> structure_batch_list = new ArrayList<>();
List<List<NDArray>> bbox_batch_list = new ArrayList<>();
List<List<NDArray>> result_score_list = new ArrayList<>();
// get ignored tokens
int beg_idx = dict.indexOf(beg_str);
int end_idx = dict.indexOf(end_str);
long batch_size = structure_idx.size(0);
for (int batch_idx = 0; batch_idx < batch_size; batch_idx++) {
List<String> structure_list = new ArrayList<>();
List<NDArray> bbox_list = new ArrayList<>();
List<NDArray> score_list = new ArrayList<>();
long len = structure_idx.get(batch_idx).size();
for (int idx = 0; idx < len; idx++) {
int char_idx = (int) structure_idx.get(batch_idx).get(idx).toLongArray()[0];
if (idx > 0 && char_idx == end_idx) {
break;
}
// if (char_idx == beg_idx || char_idx == end_idx) {
// continue;
// }
String text = dict.get(char_idx);
if(td_token.indexOf(text)>-1){
NDArray bbox = bbox_preds.get(batch_idx, idx);
// bbox.set(new NDIndex("0::2"), bbox.get(new NDIndex("0::2")));
// bbox.set(new NDIndex("1::2"), bbox.get(new NDIndex("1::2")));
bbox_list.add(bbox);
}
structure_list.add(text);
score_list.add(structure_probs.get(batch_idx, idx));
}
structure_batch_list.add(structure_list); // structure_str
bbox_batch_list.add(bbox_list);
result_score_list.add(score_list);
}
List<String> structure_str_list =structure_batch_list.get(0);
List<NDArray> bbox_list = bbox_batch_list.get(0);
List<NDArray> score_list = result_score_list.get(0);
structure_str_list.add(0,"<html>");
structure_str_list.add(1,"<body>");
structure_str_list.add(2,"<table>");
structure_str_list.add("</table>");
structure_str_list.add("</body>");
structure_str_list.add("</html>");
List<OcrItem> ocrItemList = new ArrayList<>();
for (int i = 0; i < bbox_list.size(); i++) {
NDArray box = bbox_list.get(i);
float[] arr = new float[4];
arr[0] = box.get(new NDIndex("0::2")).min().toFloatArray()[0];
arr[1] = box.get(new NDIndex("1::2")).min().toFloatArray()[0];
arr[2] = box.get(new NDIndex("0::2")).max().toFloatArray()[0];
arr[3] = box.get(new NDIndex("1::2")).max().toFloatArray()[0];
Point topLeft = new Point(arr[0] * xScale * width, arr[1] * yScale * height);
Point topRight = new Point(arr[2] * xScale * width, arr[1] * yScale * height);
Point bottomRight = new Point(arr[2] * xScale * width, arr[3] * yScale * height);
Point bottomLeft = new Point(arr[0] * xScale * width, arr[3] * yScale * height);
OcrBox ocrBox = new OcrBox(topLeft, topRight, bottomRight, bottomLeft);
//String tag = structure_str_list.get(i + 3); // 前面加了<html><body><table> 所以偏移+3
float score = score_list.get(i).toFloatArray()[0]; // 获取每个结构token的得分
OcrItem item = new OcrItem();
item.setOcrBox(ocrBox);
item.setScore(score);
//item.setTableTag(tag);
ocrItemList.add(item);
}
return new TableStructureResult(ocrItemList, structure_str_list);
}
@Override
public Batchifier getBatchifier() {
return null;
}
private NDArray ResizeTableImage(NDArray img, int height, int width, int maxLen) {
int localMax = Math.max(height, width);
float ratio = maxLen * 1.0f / localMax;
int resize_h = (int) (height * ratio);
int resize_w = (int) (width * ratio);
scale = ratio;
if(width > height){
xScale = 1f;
yScale = (float)width /(float)height;
} else{
xScale = (float)height /(float)width;
yScale = 1f;
}
img = NDImageUtils.resize(img, resize_w, resize_h);
return img;
}
private NDArray PaddingTableImage(TranslatorContext ctx, NDArray img, int maxLen) {
NDArray paddingImg = ctx.getNDManager().zeros(new Shape(maxLen, maxLen, 3), DataType.UINT8);
paddingImg.set(
new NDIndex("0:" + img.getShape().get(0) + ",0:" + img.getShape().get(1) + ",:"), img);
return paddingImg;
}
}

View File

@@ -0,0 +1,233 @@
package cn.smartjavaai.ocr.utils;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.lang3.math.NumberUtils;
import org.apache.poi.hssf.usermodel.*;
import org.apache.poi.ss.usermodel.*;
import org.apache.poi.ss.util.CellRangeAddress;
import org.dom4j.Document;
import org.dom4j.DocumentException;
import org.dom4j.DocumentHelper;
import org.dom4j.Element;
import java.util.ArrayList;
import java.util.List;
/**
* @Auther: xiaoqiang
* @Date: 2020/12/9 9:16
* @Description:
*/
public class ConvertHtml2Excel {
/**
* html表格转excel
*
* @param tableHtml 如
* <table>
* ..
* </table>
* @return
*/
public static HSSFWorkbook table2Excel(String tableHtml) {
HSSFWorkbook wb = new HSSFWorkbook();
HSSFSheet sheet = wb.createSheet();
List<CrossRangeCellMeta> crossRowEleMetaLs = new ArrayList<>();
int rowIndex = 0;
try {
Document data = DocumentHelper.parseText(tableHtml);
// 生成表头
Element thead = data.getRootElement().element("thead");
HSSFCellStyle titleStyle = getTitleStyle(wb);
int ls=0;//列数
if (thead != null) {
List<Element> trLs = thead.elements("tr");
for (Element trEle : trLs) {
HSSFRow row = sheet.createRow(rowIndex);
List<Element> thLs = trEle.elements("td");
ls=thLs.size();
makeRowCell(thLs, rowIndex, row, 0, titleStyle, crossRowEleMetaLs);
rowIndex++;
}
}
// 生成表体
Element tbody = data.getRootElement().element("tbody");
HSSFCellStyle contentStyle = getContentStyle(wb);
if (tbody != null) {
List<Element> trLs = tbody.elements("tr");
for (Element trEle : trLs) {
HSSFRow row = sheet.createRow(rowIndex);
List<Element> thLs = trEle.elements("th");
int cellIndex = makeRowCell(thLs, rowIndex, row, 0, titleStyle, crossRowEleMetaLs);
List<Element> tdLs = trEle.elements("td");
makeRowCell(tdLs, rowIndex, row, cellIndex, contentStyle, crossRowEleMetaLs);
rowIndex++;
}
}
// 合并表头
for (CrossRangeCellMeta crcm : crossRowEleMetaLs) {
sheet.addMergedRegion(new CellRangeAddress(crcm.getFirstRow(), crcm.getLastRow(), crcm.getFirstCol(), crcm.getLastCol()));
setRegionStyle(sheet, new CellRangeAddress(crcm.getFirstRow(), crcm.getLastRow(), crcm.getFirstCol(), crcm.getLastCol()),titleStyle);
}
for(int i=0;i<sheet.getRow(0).getPhysicalNumberOfCells();i++){
sheet.autoSizeColumn(i, true);//设置列宽
if(sheet.getColumnWidth(i)<255*256){
sheet.setColumnWidth(i, sheet.getColumnWidth(i) < 9000 ? 9000 : sheet.getColumnWidth(i));
}else{
sheet.setColumnWidth(i, 15000);
}
}
} catch (DocumentException e) {
e.printStackTrace();
}
return wb;
}
/**
* 生产行内容
*
* @return 最后一列的cell index
*/
/**
* @param tdLs th或者td集合
* @param rowIndex 行号
* @param row POI行对象
* @param startCellIndex
* @param cellStyle 样式
* @param crossRowEleMetaLs 跨行元数据集合
* @return
*/
private static int makeRowCell(List<Element> tdLs, int rowIndex, HSSFRow row, int startCellIndex, HSSFCellStyle cellStyle,
List<CrossRangeCellMeta> crossRowEleMetaLs) {
int i = startCellIndex;
for (int eleIndex = 0; eleIndex < tdLs.size(); i++, eleIndex++) {
int captureCellSize = getCaptureCellSize(rowIndex, i, crossRowEleMetaLs);
while (captureCellSize > 0) {
for (int j = 0; j < captureCellSize; j++) {// 当前行跨列处理(补单元格)
row.createCell(i);
i++;
}
captureCellSize = getCaptureCellSize(rowIndex, i, crossRowEleMetaLs);
}
Element thEle = tdLs.get(eleIndex);
String val = thEle.getTextTrim();
if (StringUtils.isBlank(val)) {
Element e = thEle.element("a");
if (e != null) {
val = e.getTextTrim();
}
}
HSSFCell c = row.createCell(i);
if (NumberUtils.isNumber(val)) {
c.setCellValue(Double.parseDouble(val));
c.setCellType(CellType.NUMERIC);
} else {
c.setCellValue(val);
}
int rowSpan = NumberUtils.toInt(thEle.attributeValue("rowspan"), 1);
int colSpan = NumberUtils.toInt(thEle.attributeValue("colspan"), 1);
c.setCellStyle(cellStyle);
if (rowSpan > 1 || colSpan > 1) { // 存在跨行或跨列
crossRowEleMetaLs.add(new CrossRangeCellMeta(rowIndex, i, rowSpan, colSpan));
}
if (colSpan > 1) {// 当前行跨列处理(补单元格)
for (int j = 1; j < colSpan; j++) {
i++;
row.createCell(i);
}
}
}
return i;
}
/**
* 设置合并单元格的边框样式
*
* @param sheet
* @param region
* @param cs
*/
public static void setRegionStyle(HSSFSheet sheet, CellRangeAddress region, HSSFCellStyle cs) {
for (int i = region.getFirstRow(); i <= region.getLastRow(); i++) {
HSSFRow row = sheet.getRow(i);
for (int j = region.getFirstColumn(); j <= region.getLastColumn(); j++) {
HSSFCell cell = row.getCell(j);
cell.setCellStyle(cs);
}
}
}
/**
* 获得因rowSpan占据的单元格
*
* @param rowIndex 行号
* @param colIndex 列号
* @param crossRowEleMetaLs 跨行列元数据
* @return 当前行在某列需要占据单元格
*/
private static int getCaptureCellSize(int rowIndex, int colIndex, List<CrossRangeCellMeta> crossRowEleMetaLs) {
int captureCellSize = 0;
for (CrossRangeCellMeta crossRangeCellMeta : crossRowEleMetaLs) {
if (crossRangeCellMeta.getFirstRow() < rowIndex && crossRangeCellMeta.getLastRow() >= rowIndex) {
if (crossRangeCellMeta.getFirstCol() <= colIndex && crossRangeCellMeta.getLastCol() >= colIndex) {
captureCellSize = crossRangeCellMeta.getLastCol() - colIndex + 1;
}
}
}
return captureCellSize;
}
/**
* 获得标题样式
*
* @param workbook
* @return
*/
private static HSSFCellStyle getTitleStyle(HSSFWorkbook workbook) {
//short titlebackgroundcolor = IndexedColors.GREY_25_PERCENT.index;
short fontSize = 12;
String fontName = "宋体";
HSSFCellStyle style = workbook.createCellStyle();
style.setVerticalAlignment(VerticalAlignment.CENTER);
style.setAlignment(HorizontalAlignment.CENTER);
style.setBorderBottom(BorderStyle.THIN); //下边框
style.setBorderLeft(BorderStyle.THIN);//左边框
style.setBorderTop(BorderStyle.THIN);//上边框
style.setBorderRight(BorderStyle.THIN);//右边框
//style.setFillPattern(FillPatternType.SOLID_FOREGROUND);
//style.setFillForegroundColor(titlebackgroundcolor);// 背景色
HSSFFont font = workbook.createFont();
font.setFontName(fontName);
font.setFontHeightInPoints(fontSize);
font.setBold(true);
style.setFont(font);
return style;
}
/**
* 获得内容样式
*
* @param wb
* @return
*/
private static HSSFCellStyle getContentStyle(HSSFWorkbook wb) {
short fontSize = 12;
String fontName = "宋体";
HSSFCellStyle style = wb.createCellStyle();
style.setBorderBottom(BorderStyle.THIN); //下边框
style.setBorderLeft(BorderStyle.THIN);//左边框
style.setBorderTop(BorderStyle.THIN);//上边框
style.setBorderRight(BorderStyle.THIN);//右边框
HSSFFont font = wb.createFont();
font.setFontName(fontName);
font.setFontHeightInPoints(fontSize);
style.setFont(font);
style.setAlignment(HorizontalAlignment.CENTER);//水平居中
style.setVerticalAlignment(VerticalAlignment.CENTER);//垂直居中
style.setWrapText(true);
return style;
}
}

View File

@@ -0,0 +1,42 @@
package cn.smartjavaai.ocr.utils;
/**
* @Auther: xiaoqiang
* @Date: 2020/12/9 9:17
* @Description:
*/
public class CrossRangeCellMeta {
public CrossRangeCellMeta(int firstRowIndex, int firstColIndex, int rowSpan, int colSpan) {
super();
this.firstRowIndex = firstRowIndex;
this.firstColIndex = firstColIndex;
this.rowSpan = rowSpan;
this.colSpan = colSpan;
}
private int firstRowIndex;
private int firstColIndex;
private int rowSpan;// 跨越行数
private int colSpan;// 跨越列数
public int getFirstRow() {
return firstRowIndex;
}
public int getLastRow() {
return firstRowIndex + rowSpan - 1;
}
public int getFirstCol() {
return firstColIndex;
}
public int getLastCol() {
return firstColIndex + colSpan - 1;
}
public int getColSpan(){
return colSpan;
}
}

View File

@@ -11,14 +11,12 @@ import ai.djl.ndarray.NDManager;
import ai.djl.opencv.OpenCVImageFactory;
import cn.smartjavaai.common.entity.*;
import cn.smartjavaai.common.entity.Point;
import cn.smartjavaai.ocr.entity.OcrBox;
import cn.smartjavaai.ocr.entity.OcrInfo;
import cn.smartjavaai.ocr.entity.OcrItem;
import cn.smartjavaai.ocr.entity.RotatedBoxCompX;
import cn.smartjavaai.ocr.entity.*;
import cn.smartjavaai.ocr.enums.AngleEnum;
import cn.smartjavaai.ocr.opencv.OcrNDArrayUtils;
import cn.smartjavaai.ocr.opencv.OcrOpenCVUtils;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.opencv.core.Mat;
import org.opencv.core.Scalar;
import org.opencv.imgproc.Imgproc;
@@ -26,10 +24,8 @@ import org.opencv.imgproc.Imgproc;
import java.awt.*;
import java.awt.image.BufferedImage;
import java.math.BigDecimal;
import java.util.ArrayList;
import java.util.Iterator;
import java.util.*;
import java.util.List;
import java.util.Objects;
/**
* @author dwj
@@ -42,27 +38,41 @@ public class OcrUtils {
/**
* 转换为OcrBox
* @param dt_boxes
* @param img
* @return
*/
public static List<OcrBox> convertToOcrBox(NDList dt_boxes, Image img){
if(Objects.isNull(dt_boxes) || dt_boxes.size() == 0){
return null;
}
List<OcrBox> boxList = new ArrayList<OcrBox>();
for(NDArray box : dt_boxes){
public static List<OcrBox> convertToOcrBox(NDList dt_boxes) {
List<OcrBox> boxList = new ArrayList<>();
for (NDArray box : dt_boxes) {
float[] pointsArr = box.toFloatArray();
//log.debug("points: {}", pointsArr);
float[] lt = java.util.Arrays.copyOfRange(pointsArr, 0, 2);
float[] rt = java.util.Arrays.copyOfRange(pointsArr, 2, 4);
float[] rb = java.util.Arrays.copyOfRange(pointsArr, 4, 6);
float[] lb = java.util.Arrays.copyOfRange(pointsArr, 6, 8);
OcrBox ocrBox = new OcrBox(new Point(lt[0], lt[1]), new Point(rt[0], rt[1]), new Point(rb[0], rb[1]), new Point(lb[0], lb[1]));
OcrBox ocrBox = new OcrBox(
new Point(pointsArr[0], pointsArr[1]),
new Point(pointsArr[2], pointsArr[3]),
new Point(pointsArr[4], pointsArr[5]),
new Point(pointsArr[6], pointsArr[7])
);
boxList.add(ocrBox);
}
return boxList;
}
/**
* 转换为OcrBox
* @param dt_boxes
* @return
*/
public static List<List<OcrBox>> convertToOcrBox(List<NDList> ndLists) {
if (ndLists == null || ndLists.isEmpty()) {
return Collections.emptyList();
}
List<List<OcrBox>> boxLists = new ArrayList<>();
for (NDList dt_boxes : ndLists) {
boxLists.add(convertToOcrBox(dt_boxes));
}
return boxLists;
}
/**
* 欧式距离计算
*
@@ -140,7 +150,6 @@ public class OcrUtils {
if(Objects.isNull(lines) || lines.size() == 0){
return null;
}
List<DetectionInfo> detectionInfoList = new ArrayList<DetectionInfo>();
List<List<OcrItem>> lineList = new ArrayList<List<OcrItem>>();
String fullText = "";
for(ArrayList<RotatedBoxCompX> boxList : lines){
@@ -165,6 +174,36 @@ public class OcrUtils {
return new OcrInfo(lineList, fullText);
}
public static OcrInfo convertRotatedBoxesToOcrItems(List<RotatedBox> rotatedBoxes) {
OcrInfo ocrInfo = new OcrInfo();
List<OcrItem> ocrItems = new ArrayList<>();
StringBuilder fullText = new StringBuilder();
for (RotatedBox rotatedBox : rotatedBoxes) {
NDArray box = rotatedBox.getBox();
float[] points = box.toFloatArray();
Point topLeft = new Point(points[0], points[1]);
Point topRight = new Point(points[2], points[3]);
Point bottomRight = new Point(points[4], points[5]);
Point bottomLeft = new Point(points[6], points[7]);
OcrBox ocrBox = new OcrBox(topLeft, topRight, bottomRight, bottomLeft);
String text = rotatedBox.getText();
OcrItem item = new OcrItem();
item.setOcrBox(ocrBox);
item.setText(text);
ocrItems.add(item);
fullText.append(text + " ");
}
if (fullText.length() > 0) {
fullText.deleteCharAt(fullText.length() - 1);
}
ocrInfo.setOcrItemList(ocrItems);
ocrInfo.setFullText(fullText.toString());
return ocrInfo;
}
/**
* 放射变换+裁剪
@@ -235,26 +274,28 @@ public class OcrUtils {
// 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角
BasicStroke bStroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER);
g.setStroke(bStroke);
for(List<OcrItem> ocrItemList : ocrInfo.getLineList()){
for(OcrItem item : ocrItemList){
OcrBox box = item.getOcrBox();
int[] xPoints = {
(int)box.getTopLeft().getX(),
(int)box.getTopRight().getX(),
(int)box.getBottomRight().getX(),
(int)box.getBottomLeft().getX(),
(int)box.getTopLeft().getX()
};
int[] yPoints = {
(int)box.getTopLeft().getY(),
(int)box.getTopRight().getY(),
(int)box.getBottomRight().getY(),
(int)box.getBottomLeft().getY(),
(int)box.getTopLeft().getY()
};
g.drawPolyline(xPoints, yPoints, 5);
g.drawString(item.getText(), xPoints[0], yPoints[0]);
}
List<OcrItem> ocrItemList = ocrInfo.getOcrItemList();
if(CollectionUtils.isNotEmpty(ocrInfo.getLineList())){
ocrItemList = ocrInfo.flattenLines();
}
for(OcrItem item : ocrItemList){
OcrBox box = item.getOcrBox();
int[] xPoints = {
(int)box.getTopLeft().getX(),
(int)box.getTopRight().getX(),
(int)box.getBottomRight().getX(),
(int)box.getBottomLeft().getX(),
(int)box.getTopLeft().getX()
};
int[] yPoints = {
(int)box.getTopLeft().getY(),
(int)box.getTopRight().getY(),
(int)box.getBottomRight().getY(),
(int)box.getBottomLeft().getY(),
(int)box.getTopLeft().getY()
};
g.drawPolyline(xPoints, yPoints, 5);
g.drawString(item.getText(), xPoints[0], yPoints[0]);
}
} finally {
g.dispose();