From c8cda3f24070b39e91d7c5caf5a5fe162fe5f093 Mon Sep 17 00:00:00 2001 From: dengwenjie <775747758@qq.com> Date: Sun, 29 Mar 2026 15:08:02 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96OCR=E8=AF=86=E5=88=AB?= =?UTF-8?q?=E5=86=85=E5=AD=98=E7=AE=A1=E7=90=86=E5=B9=B6=E6=94=AF=E6=8C=81?= =?UTF-8?q?=E5=88=86=E5=9D=97=E8=AF=86=E5=88=AB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../smartjavaai/ocr/config/OcrRecOptions.java | 1 + .../common/detect/OcrCommonDetModelImpl.java | 15 +- .../direction/PPOCRMobileV2ClsModel.java | 17 +- .../common/recognize/OcrCommonRecModel.java | 25 ++ .../recognize/OcrCommonRecModelImpl.java | 197 +++++++++---- .../cn/smartjavaai/ocr/utils/BoxUtils.java | 278 ++++++++++++++++++ .../cn/smartjavaai/ocr/utils/OcrUtils.java | 39 +++ 7 files changed, 506 insertions(+), 66 deletions(-) create mode 100644 ocr/src/main/java/cn/smartjavaai/ocr/utils/BoxUtils.java diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java index 310e4a6..ab9921e 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java @@ -21,6 +21,7 @@ public class OcrRecOptions { private boolean enableLineSplit = true; + public OcrRecOptions(boolean enableDirectionCorrect, boolean enableLineSplit) { this.enableDirectionCorrect = enableDirectionCorrect; this.enableLineSplit = enableLineSplit; diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java index 5a2e828..efff96e 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java @@ -94,8 +94,10 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ @Override public List detect(Image image){ + long start = System.nanoTime(); List imageList = Collections.singletonList(image); List> result = batchDetectDJLImage(imageList); + log.debug("文本检测模型单图调用耗时={}ms, 检测框数量={}", elapsedMillis(start), result.isEmpty() ? 0 : result.get(0).size()); return result.get(0); } @@ -192,12 +194,19 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ if(!ImageUtils.isAllImageSizeEqual(imageList)){ throw new OcrException("图片尺寸不一致"); } + long totalStart = System.nanoTime(); Predictor predictor = null; try (NDManager manager = NDManager.newBaseManager()) { predictor = detPredictorPool.borrowObject(); List result = predictor.batchPredict(imageList); result.forEach(ndList -> ndList.attach(manager)); - return OcrUtils.convertToOcrBox(result); + List> boxes = OcrUtils.convertToOcrBox(result); + int totalBoxes = 0; + for (List boxList : boxes) { + totalBoxes += boxList == null ? 0 : boxList.size(); + } + log.debug("文本检测总耗时={}ms, batchSize={}, totalBoxes={}", elapsedMillis(totalStart), imageList.size(), totalBoxes); + return boxes; } catch (Exception e) { throw new OcrException("OCR检测错误", e); }finally { @@ -257,5 +266,9 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ return fromFactory; } + private long elapsedMillis(long startNanos) { + return (System.nanoTime() - startNanos) / 1_000_000; + } + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java index 2b70eb1..0acdbd2 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java @@ -108,13 +108,13 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { if(Objects.isNull(textDetModel)){ throw new OcrException("textDetModel is null"); } - //检测文本 List boxeList = textDetModel.detect(image); if(Objects.isNull(boxeList) || boxeList.isEmpty()){ throw new OcrException("未检测到文本"); } Mat srcMat = ImageUtils.toMat(image); - return detect(boxeList, srcMat); + List result = detect(boxeList, srcMat); + return result; } @@ -274,7 +274,9 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { List imageList = new ArrayList(); List isRotatedList = new ArrayList(); int index = 0; + long totalStart = System.nanoTime(); try (NDManager manager = model.getNDManager().newSubManager()){ + long prepareStart = System.nanoTime(); for(int i = 0; i < srcMatList.size(); i++){ for (int j = 0; j < boxList.get(i).size(); j++){ //透视变换及裁剪 @@ -292,8 +294,11 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { index++; } } + log.debug("方向模型裁剪预处理耗时={}ms, batchSize={}, textBlocks={}", elapsedMillis(prepareStart), srcMatList.size(), imageList.size()); List> result = new ArrayList<>(); + long predictStart = System.nanoTime(); List directionInfos = batchDetect(imageList); + log.debug("方向分类模型调用耗时={}ms, textBlocks={}", elapsedMillis(predictStart), imageList.size()); //释放 imageList.forEach(image -> ImageUtils.releaseOpenCVMat(image)); if(CollectionUtils.isEmpty(directionInfos)){ @@ -327,6 +332,7 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { } result.add(ocrItemList); } + log.debug("方向模型总耗时={}ms, batchSize={}, textBlocks={}", elapsedMillis(totalStart), srcMatList.size(), index); return result; } } @@ -335,7 +341,8 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { Predictor predictor = null; try { predictor = predictorPool.borrowObject(); - return predictor.batchPredict(imageList); + List result = predictor.batchPredict(imageList); + return result; } catch (Exception e) { throw new OcrException("OCR检测错误", e); }finally { @@ -399,4 +406,8 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { public boolean isFromFactory() { return fromFactory; } + + private long elapsedMillis(long startNanos) { + return (System.nanoTime() - startNanos) / 1_000_000; + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java index 98e90ed..3966e3e 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java @@ -61,6 +61,19 @@ public interface OcrCommonRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 基于已有文本检测框执行文本识别。 + * 适合调用方已经完成文本检测,想避免重复检测的场景。 + * + * @param image 原图 + * @param boxList 已有文本检测框 + * @param options 识别选项 + * @return OCR 结果 + */ + default OcrInfo recognize(Image image, List boxList, OcrRecOptions options) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 文本检测 @@ -141,6 +154,18 @@ public interface OcrCommonRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 基于已有文本检测框执行批量文本识别。 + * + * @param imageList 原图列表 + * @param boxList 每张图对应的文本检测框列表 + * @param options 识别选项 + * @return OCR 结果 + */ + default List batchRecognizeDJLImage(List imageList, List> boxList, OcrRecOptions options) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + default GenericObjectPool> getPool() { throw new UnsupportedOperationException("默认不支持该功能"); } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java index 7d9eeb4..03f6e86 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java @@ -47,6 +47,8 @@ import java.util.stream.Collectors; @Slf4j public class OcrCommonRecModelImpl implements OcrCommonRecModel { + private static final int REC_CHUNK_SIZE = 64; + private GenericObjectPool> recPredictorPool; private OcrRecModelConfig config; @@ -111,10 +113,27 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { */ @Override public OcrInfo recognize(Image image, OcrRecOptions options) { + long start = System.nanoTime(); List result = batchRecognizeDJLImage(Collections.singletonList(image), options); if (CollectionUtils.isEmpty(result)) { throw new OcrException("OCR识别结果为空"); } + log.debug("OCR识别单图总耗时={}ms", elapsedMillis(start)); + return result.get(0); + } + + @Override + public OcrInfo recognize(Image image, List boxList, OcrRecOptions options) { + long start = System.nanoTime(); + List result = batchRecognizeDJLImage( + Collections.singletonList(image), + Collections.singletonList(boxList), + options + ); + if (CollectionUtils.isEmpty(result)) { + throw new OcrException("OCR识别结果为空"); + } + log.debug("OCR识别单图总耗时={}ms", elapsedMillis(start)); return result.get(0); } @@ -212,10 +231,8 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { //分行判断 for (int i = 1; i < initList.size(); i++) { RotatedBoxCompX tmpBox = new RotatedBoxCompX(initList.get(i).getBox(), initList.get(i).getText()); - float y1 = firstBox.getBox().toFloatArray()[1]; - float y2 = tmpBox.getBox().toFloatArray()[1]; - float dis = Math.abs(y2 - y1); - if (dis < 20) { // 认为是同 1 行 - Considered to be in the same line + boolean isSameRow = OcrUtils.isSameRow(firstBox.getBox(), tmpBox.getBox()); + if (isSameRow) { line.add(tmpBox); } else { // 换行 - Line break firstBox = tmpBox; @@ -346,6 +363,11 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { @Override public List batchRecognizeDJLImage(List imageList, OcrRecOptions options) { + return batchRecognizeDJLImage(imageList, null, options); + } + + @Override + public List batchRecognizeDJLImage(List imageList, List> boxeList, OcrRecOptions options) { if (Objects.isNull(textDetModel)) { throw new OcrException("textDetModel is null"); } @@ -356,64 +378,82 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { if (CollectionUtils.isEmpty(imageList)) { throw new OcrException("imageList is empty"); } - //检测文本 - List> boxeList = textDetModel.batchDetectDJLImage(imageList); - if (CollectionUtils.isEmpty(boxeList) || boxeList.size() != imageList.size()) { + long totalStart = System.nanoTime(); + List> effectiveBoxList = boxeList; + if (CollectionUtils.isEmpty(effectiveBoxList)) { + effectiveBoxList = textDetModel.batchDetectDJLImage(imageList); + } + if (CollectionUtils.isEmpty(effectiveBoxList) || effectiveBoxList.size() != imageList.size()) { throw new OcrException("未检测到文本"); } Predictor predictor = null; List ocrInfoList = new ArrayList(); try (NDManager manager = NDManager.newBaseManager()) { predictor = recPredictorPool.borrowObject(); - List allImageAlignList = new ArrayList(); //检测方向 if (ocrRecOptions.isEnableDirectionCorrect()) { if (Objects.isNull(directionModel)) { throw new OcrException("请配置方向模型"); } - List matList = imageList.stream() - .map(image -> ImageUtils.toMat(image)) - .collect(Collectors.toList()); - List> ocrItemList = directionModel.batchDetect(boxeList, matList); - if (CollectionUtils.isEmpty(ocrItemList) || ocrItemList.size() != imageList.size()) { - throw new OcrException("方向检测失败"); - } - allImageAlignList = new ArrayList(); - for (int i = 0; i < ocrItemList.size(); i++) { - Mat srcMat = ImageUtils.toMat(imageList.get(i)); - List 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); + long directionStart = System.nanoTime(); + List matList = new ArrayList<>(imageList.size()); + try { + for (Image image : imageList) { + matList.add(ImageUtils.toMat(image)); + } + List> ocrItemList = directionModel.batchDetect(effectiveBoxList, matList); + log.debug("OCR流程-文本方向分类耗时={}ms, batchSize={}", elapsedMillis(directionStart), imageList.size()); + if (CollectionUtils.isEmpty(ocrItemList) || ocrItemList.size() != imageList.size()) { + throw new OcrException("方向检测失败"); + } + long alignStart = System.nanoTime(); + List textList = new ArrayList<>(); + List chunkImages = new ArrayList<>(REC_CHUNK_SIZE); + for (int i = 0; i < ocrItemList.size(); i++) { + Mat srcMat = matList.get(i); + List imageAlignList = batchAlignWithDirection(ocrItemList.get(i), srcMat, manager); + for (Image alignImage : imageAlignList) { + chunkImages.add(alignImage); + if (chunkImages.size() >= REC_CHUNK_SIZE) { + textList.addAll(batchRecognizeChunk(predictor, chunkImages)); + } + } + } + log.debug("OCR流程-方向矫正裁剪耗时={}ms, textBlocks={}", elapsedMillis(alignStart), textList.size() + chunkImages.size()); + long recStart = System.nanoTime(); + if (!chunkImages.isEmpty()) { + textList.addAll(batchRecognizeChunk(predictor, chunkImages)); + } + log.debug("OCR流程-识别模型调用耗时={}ms, textBlocks={}", elapsedMillis(recStart), textList.size()); + return buildOcrInfoList(effectiveBoxList, ocrRecOptions, manager, textList, imageList.size(), totalStart); + } finally { + releaseTemporaryMats(imageList, matList); } } else { - for (int i = 0; i < boxeList.size(); i++) { - Mat srcMat = ImageUtils.toMat(imageList.get(i)); - List imageAlignList = batchAlign(boxeList.get(i), srcMat, manager); -// for(int j = 0; j < imageAlignList.size(); j++){ -// ImageUtils.saveImage(imageAlignList.get(j),i+"-"+j+".png","/Users/wenjie/Downloads/testing33"); -// } - allImageAlignList.addAll(imageAlignList); - } - } - List textList = batchRecognize(allImageAlignList); - int textIndex = 0; - for (int i = 0; i < boxeList.size(); i++) { - List rotatedBoxes = new ArrayList<>(); - for (int j = 0; j < boxeList.get(i).size(); j++) { - if (textIndex >= textList.size()) { - throw new OcrException("识别失败: 第" + i + "张图片, 第" + j + "个文本块,未识别到文本"); + List textList = new ArrayList<>(); + List chunkImages = new ArrayList<>(REC_CHUNK_SIZE); + for (int i = 0; i < effectiveBoxList.size(); i++) { + Mat srcMat = null; + try { + srcMat = ImageUtils.toMat(imageList.get(i)); + List imageAlignList = batchAlign(effectiveBoxList.get(i), srcMat, manager); + for (Image alignImage : imageAlignList) { + chunkImages.add(alignImage); + if (chunkImages.size() >= REC_CHUNK_SIZE) { + textList.addAll(batchRecognizeChunk(predictor, chunkImages)); + } + } + } finally { + releaseTemporaryMat(imageList.get(i), srcMat); } - 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); + long recStart = System.nanoTime(); + if (!chunkImages.isEmpty()) { + textList.addAll(batchRecognizeChunk(predictor, chunkImages)); + } + log.debug("OCR流程-识别模型调用耗时={}ms, textBlocks={}", elapsedMillis(recStart), textList.size()); + return buildOcrInfoList(effectiveBoxList, ocrRecOptions, manager, textList, imageList.size(), totalStart); } - return ocrInfoList; } catch (Exception e) { throw new OcrException("OCR检测错误", e); } finally { @@ -432,28 +472,41 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { } } - private List batchRecognize(List imageAlignList) { - Predictor predictor = null; + private List buildOcrInfoList(List> effectiveBoxList, + OcrRecOptions ocrRecOptions, + NDManager manager, + List textList, + int batchSize, + long totalStart) { + int textIndex = 0; + List ocrInfoList = new ArrayList<>(); + for (int i = 0; i < effectiveBoxList.size(); i++) { + List rotatedBoxes = new ArrayList<>(); + for (int j = 0; j < effectiveBoxList.get(i).size(); j++) { + if (textIndex >= textList.size()) { + throw new OcrException("识别失败: 第" + i + "张图片, 第" + j + "个文本块,未识别到文本"); + } + OcrBox box = effectiveBoxList.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); + } + log.debug("OCR流程总耗时={}ms, batchSize={}", elapsedMillis(totalStart), batchSize); + return ocrInfoList; + } + + private List batchRecognizeChunk(Predictor predictor, List imageAlignList) { try { - predictor = recPredictorPool.borrowObject(); List textList = predictor.batchPredict(imageAlignList); - imageAlignList.forEach(subImg -> ImageUtils.releaseOpenCVMat(subImg)); 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); - } - } - } + imageAlignList.forEach(ImageUtils::releaseOpenCVMat); + imageAlignList.clear(); } } @@ -525,4 +578,24 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { public boolean isFromFactory() { return fromFactory; } + + private long elapsedMillis(long startNanos) { + return (System.nanoTime() - startNanos) / 1_000_000; + } + + private void releaseTemporaryMats(List imageList, List matList) { + int size = Math.min(imageList.size(), matList.size()); + for (int i = 0; i < size; i++) { + releaseTemporaryMat(imageList.get(i), matList.get(i)); + } + } + + private void releaseTemporaryMat(Image image, Mat mat) { + if (mat == null || image == null) { + return; + } + if (!(image.getWrappedImage() instanceof Mat)) { + mat.release(); + } + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/utils/BoxUtils.java b/ocr/src/main/java/cn/smartjavaai/ocr/utils/BoxUtils.java new file mode 100644 index 0000000..7446c8e --- /dev/null +++ b/ocr/src/main/java/cn/smartjavaai/ocr/utils/BoxUtils.java @@ -0,0 +1,278 @@ +package cn.smartjavaai.ocr.utils; + +import cn.smartjavaai.common.entity.Point; +import cn.smartjavaai.ocr.entity.OcrBox; + +import java.util.ArrayList; +import java.util.Comparator; +import java.util.List; + +/** + * @author dwj + * @date 2026/1/11 + */ +public class BoxUtils { + + + public enum Direction { + UP, DOWN, LEFT, RIGHT + } + + /** + * 寻找指定方向上距离最近的框 + * + * @param anchor 锚点框,作为搜索的起始参考框 + * @param boxList 候选框列表,在其中搜索目标框 + * @param direction 搜索方向,可选 UP(上)、DOWN(下)、LEFT(左)、RIGHT(右) + * @return 找到的最近的 OcrBox,如果找不到符合条件的框则返回 null + * + * @apiNote + * - 该方法会在指定方向上寻找与锚点框在同一行或同一列的最近邻框 + * - 水平方向(LEFT/RIGHT):要求候选框与锚点框在 Y 轴上有重叠(同一行) + * - 垂直方向(UP/DOWN):要求候选框与锚点框在 X 轴上有重叠(同一列) + * - 使用欧几里得距离计算中心点之间的距离 + */ + public static OcrBox findNearestBox(OcrBox anchor, List boxList, Direction direction) { + OcrBox nearest = null; + double bestScore = Double.MAX_VALUE; + AxisSystem axisSystem = buildAxisSystem(anchor); + + for (OcrBox target : boxList) { + if (target == anchor) { + continue; + } + + CandidateMetrics metrics = evaluateCandidate(anchor, target, direction, axisSystem); + if (metrics == null) { + continue; + } + if (metrics.score < bestScore) { + bestScore = metrics.score; + nearest = target; + } + } + + return nearest; + } + + /** + * 寻找指定方向上距离最近的多个框(按距离升序返回) + * + * @param anchor 锚点框 + * @param boxList 候选框列表 + * @param direction 搜索方向 + * @param limit 返回的最大数量(<=0 时返回空列表) + * @return 按距离由近到远排序的 OcrBox 列表 + */ + public static List findNearestBoxes(OcrBox anchor, List boxList, Direction direction, int limit) { + if (anchor == null || boxList == null || boxList.isEmpty() || limit <= 0) { + return new ArrayList<>(); + } + + AxisSystem axisSystem = buildAxisSystem(anchor); + List candidates = new ArrayList<>(); + + for (OcrBox target : boxList) { + if (target == anchor) { + continue; + } + + CandidateMetrics metrics = evaluateCandidate(anchor, target, direction, axisSystem); + if (metrics == null) { + continue; + } + candidates.add(new Neighbor(target, metrics.score)); + } + + candidates.sort(Comparator.comparingDouble(n -> n.distance)); + + List result = new ArrayList<>(); + int size = Math.min(limit, candidates.size()); + for (int i = 0; i < size; i++) { + result.add(candidates.get(i).box); + } + return result; + } + + /** + * 内部使用的邻居结构体,存储框和距离 + */ + private static class Neighbor { + private final OcrBox box; + private final double distance; + + private Neighbor(OcrBox box, double distance) { + this.box = box; + this.distance = distance; + } + } + + private static CandidateMetrics evaluateCandidate(OcrBox anchor, OcrBox target, Direction direction, AxisSystem axisSystem) { + double mainAxisX = isHorizontalDirection(direction) ? axisSystem.horizontalAxisX : axisSystem.verticalAxisX; + double mainAxisY = isHorizontalDirection(direction) ? axisSystem.horizontalAxisY : axisSystem.verticalAxisY; + double crossAxisX = isHorizontalDirection(direction) ? axisSystem.verticalAxisX : axisSystem.horizontalAxisX; + double crossAxisY = isHorizontalDirection(direction) ? axisSystem.verticalAxisY : axisSystem.horizontalAxisY; + + Projection anchorMain = projectBox(anchor, mainAxisX, mainAxisY); + Projection anchorCross = projectBox(anchor, crossAxisX, crossAxisY); + Projection targetMain = projectBox(target, mainAxisX, mainAxisY); + Projection targetCross = projectBox(target, crossAxisX, crossAxisY); + Point anchorCenter = getCenter(anchor); + Point targetCenter = getCenter(target); + double centerMainDelta = projectPointDelta(anchorCenter, targetCenter, mainAxisX, mainAxisY); + + double mainGap; + switch (direction) { + case RIGHT: + case DOWN: + mainGap = targetMain.min - anchorMain.max; + if (centerMainDelta <= 0) { + return null; + } + break; + case LEFT: + case UP: + mainGap = anchorMain.min - targetMain.max; + if (centerMainDelta >= 0) { + return null; + } + break; + default: + return null; + } + + double anchorMainSize = Math.max(1.0, anchorMain.max - anchorMain.min); + double targetMainSize = Math.max(1.0, targetMain.max - targetMain.min); + double allowedBacktrack = Math.min(anchorMainSize, targetMainSize) * 0.35; + if (mainGap < -allowedBacktrack) { + return null; + } + + double overlap = Math.max(0.0, Math.min(anchorCross.max, targetCross.max) - Math.max(anchorCross.min, targetCross.min)); + double minCrossSize = Math.max(1.0, Math.min(anchorCross.max - anchorCross.min, targetCross.max - targetCross.min)); + double overlapRatio = overlap / minCrossSize; + + double anchorCrossCenter = (anchorCross.min + anchorCross.max) / 2.0; + double targetCrossCenter = (targetCross.min + targetCross.max) / 2.0; + double crossCenterDistance = Math.abs(anchorCrossCenter - targetCrossCenter); + double crossTolerance = Math.max(anchorCross.max - anchorCross.min, targetCross.max - targetCross.min) * 0.6; + + if (overlapRatio < 0.2 && crossCenterDistance > crossTolerance) { + return null; + } + + double score = Math.max(0.0, mainGap) * 10.0 + crossCenterDistance + Math.abs(centerMainDelta) * 0.01; + if (overlapRatio < 0.2) { + score += (0.2 - overlapRatio) * 100.0; + } + return new CandidateMetrics(score); + } + + private static boolean isHorizontalDirection(Direction direction) { + return direction == Direction.LEFT || direction == Direction.RIGHT; + } + + private static Point getCenter(OcrBox box) { + float cx = (float) (box.getTopLeft().getX() + box.getTopRight().getX() + box.getBottomRight().getX() + box.getBottomLeft().getX()) / 4; + float cy = (float) (box.getTopLeft().getY() + box.getTopRight().getY() + box.getBottomRight().getY() + box.getBottomLeft().getY()) / 4; + return new Point(cx, cy); + } + + private static AxisSystem buildAxisSystem(OcrBox anchor) { + Point topLeft = anchor.getTopLeft(); + Point topRight = anchor.getTopRight(); + Point bottomLeft = anchor.getBottomLeft(); + + double horizontalX = topRight.getX() - topLeft.getX(); + double horizontalY = topRight.getY() - topLeft.getY(); + double verticalX = bottomLeft.getX() - topLeft.getX(); + double verticalY = bottomLeft.getY() - topLeft.getY(); + + double horizontalNorm = Math.hypot(horizontalX, horizontalY); + double verticalNorm = Math.hypot(verticalX, verticalY); + + if (horizontalNorm < 1e-6) { + horizontalX = 1.0; + horizontalY = 0.0; + horizontalNorm = 1.0; + } + if (verticalNorm < 1e-6) { + verticalX = 0.0; + verticalY = 1.0; + verticalNorm = 1.0; + } + + return new AxisSystem( + horizontalX / horizontalNorm, + horizontalY / horizontalNorm, + verticalX / verticalNorm, + verticalY / verticalNorm + ); + } + + private static Projection projectBox(OcrBox box, double axisX, double axisY) { + double[] values = new double[]{ + dot(box.getTopLeft(), axisX, axisY), + dot(box.getTopRight(), axisX, axisY), + dot(box.getBottomRight(), axisX, axisY), + dot(box.getBottomLeft(), axisX, axisY) + }; + double min = values[0]; + double max = values[0]; + for (int i = 1; i < values.length; i++) { + min = Math.min(min, values[i]); + max = Math.max(max, values[i]); + } + return new Projection(min, max); + } + + private static double dot(Point point, double axisX, double axisY) { + return point.getX() * axisX + point.getY() * axisY; + } + + private static double projectPointDelta(Point from, Point to, double axisX, double axisY) { + return (to.getX() - from.getX()) * axisX + (to.getY() - from.getY()) * axisY; + } + + private static double min(double a, double b, double c, double d) { + return Math.min(Math.min(a, b), Math.min(c, d)); + } + + private static double max(double a, double b, double c, double d) { + return Math.max(Math.max(a, b), Math.max(c, d)); + } + + private static class AxisSystem { + private final double horizontalAxisX; + private final double horizontalAxisY; + private final double verticalAxisX; + private final double verticalAxisY; + + private AxisSystem(double horizontalAxisX, double horizontalAxisY, double verticalAxisX, double verticalAxisY) { + this.horizontalAxisX = horizontalAxisX; + this.horizontalAxisY = horizontalAxisY; + this.verticalAxisX = verticalAxisX; + this.verticalAxisY = verticalAxisY; + } + } + + private static class Projection { + private final double min; + private final double max; + + private Projection(double min, double max) { + this.min = min; + this.max = max; + } + } + + private static class CandidateMetrics { + private final double score; + + private CandidateMetrics(double score) { + this.score = score; + } + } + + +} diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java b/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java index 83dc2b1..2b9a000 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java @@ -8,6 +8,7 @@ import ai.djl.modality.cv.util.NDImageUtils; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.Shape; import ai.djl.opencv.OpenCVImageFactory; import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; @@ -412,4 +413,42 @@ public class OcrUtils { } + public static boolean isSameRow(NDArray box1, NDArray box2) { + // 1. 确保 box 是 (4, 2) 的形状 + NDArray b1 = formatBox(box1); + NDArray b2 = formatBox(box2); + + // 2. 获取 Y 坐标列(索引为 1 的列) + NDArray y1 = b1.get(":, 1"); + NDArray y2 = b2.get(":, 1"); + + float yMin1 = y1.min().getFloat(); + float yMax1 = y1.max().getFloat(); + float yMin2 = y2.min().getFloat(); + float yMax2 = y2.max().getFloat(); + + // 3. 计算重叠高度 + float overlapHeight = Math.min(yMax1, yMax2) - Math.max(yMin1, yMin2); + + if (overlapHeight <= 0) return false; + + // 4. 计算各自高度 + float h1 = yMax1 - yMin1; + float h2 = yMax2 - yMin2; + + // 5. 判定标准 + return overlapHeight > (Math.min(h1, h2) * 0.5f); + } + + /** + * 辅助方法:将 1D 的 8个元素 转换为 2D 的 (4, 2) + */ + private static NDArray formatBox(NDArray box) { + if (box.getShape().dimension() == 1) { + // 如果是 [x0, y0, x1, y1...] 这种 8 个元素的 1D 阵,转为 (4, 2) + return box.reshape(new Shape(4, 2)); + } + return box; + } + }